hcm21 / demo.py
ParetoOptimal's picture
Add demo script, API endpoints, and training notebook baselines
80d9920
Raw
History Blame Contribute Delete
30.3 kB
"""Standalone demo: run a full 6-quarter HR simulation with a heuristic strategy.
No server, no LLM API key required. Seeds the environment with sample data and
executes data-driven HR decisions through all HCM:21 phases each quarter.
Usage:
python demo.py # Default: seed=42, size=300
python demo.py --seed 123 --size 250 --scenario budget_cuts
python demo.py --scenario high_eng_turnover
python demo.py --all-scenarios # Run all 4 task scenarios
API usage (from server):
from demo import run_demo_json
result = run_demo_json(seed=42, size=300, scenario="high_eng_turnover")
"""
from __future__ import annotations
import argparse
import sys
from typing import Any, Dict, List, Optional
from hr_env.models import HRAction
from hr_env.server.environment import HRProductivityEnvironment
DEPARTMENTS = ["Engineering", "Sales", "Operations", "HR", "Finance"]
SCENARIOS = [
{"name": "high_eng_turnover", "seed": 42, "size": 300},
{"name": "budget_cuts", "seed": 123, "size": 250},
{"name": "rapid_growth", "seed": 456, "size": 350},
{"name": "balanced_optimization", "seed": 789, "size": 300},
]
# ── Formatting Helpers ──────────────────────────────────────────────
def fmt_dollar(v: float) -> str:
if abs(v) >= 1_000_000:
return f"${v / 1_000_000:,.1f}M"
if abs(v) >= 1_000:
return f"${v:,.0f}"
return f"${v:.2f}"
def fmt_pct(curr: float, prev: float) -> str:
if prev == 0:
return " N/A"
change = (curr - prev) / abs(prev) * 100
return f"{change:+.1f}%"
def print_header(seed: int, size: int, scenario: Optional[str]) -> None:
print()
print("=" * 64)
print(" HCM:21 HR Productivity Environment - Demo Playthrough")
print("=" * 64)
scn = scenario or "default (no scenario)"
print(f" Seed: {seed} | Company Size: {size} | Scenario: {scn}")
print()
def print_baseline(obs_data: dict) -> None:
summary = obs_data.get("company_summary", {})
metrics = obs_data.get("baseline_metrics", {})
print("--- COMPANY BASELINE ---")
print(f" Employees: {summary.get('total_headcount', '?')} across 5 departments")
print(f" Revenue: {fmt_dollar(summary.get('revenue', 0))} | Profit: {fmt_dollar(summary.get('profit', 0))}")
hcva = metrics.get("hcva", 0)
hcroi = metrics.get("hcroi", 0)
qips_c = metrics.get("qips", {}).get("composite", 0)
print(f" HCVA: {fmt_dollar(hcva)}/FTE | HCROI: {hcroi:.2f}x | QIPS: {qips_c:.4f}")
# Department breakdown
depts = summary.get("departments", {})
for name in DEPARTMENTS:
d = depts.get(name, {})
print(
f" {name:12s}: {d.get('headcount', '?'):>3} employees, "
f"perf {d.get('avg_performance', 0):.1f}, "
f"engage {d.get('avg_engagement', 0):.0f}, "
f"risk {d.get('avg_flight_risk', 0):.2f}"
)
print()
def print_quarter_header(q: int) -> None:
print(f"{'=' * 64}")
print(f" QUARTER {q}")
print(f"{'=' * 64}")
def print_phase(phase: str) -> None:
print(f"\n --- {phase.upper()} ---")
def print_metrics_table(
metrics: dict, prev_metrics: Optional[dict], reward: Optional[float]
) -> None:
prev = prev_metrics or {}
hcva = metrics.get("hcva", 0)
hcroi = metrics.get("hcroi", 0)
qips = metrics.get("qips", {}).get("composite", 0)
ev = metrics.get("employee_value", 0)
snap = metrics.get("snapshot", {})
headcount = snap.get("headcount", 0)
engagement = snap.get("avg_engagement", 0)
p_hcva = prev.get("hcva", 0)
p_hcroi = prev.get("hcroi", 0)
p_qips = prev.get("qips", {}).get("composite", 0)
p_snap = prev.get("snapshot", {})
p_hc = p_snap.get("headcount", 0)
p_eng = p_snap.get("avg_engagement", 0)
print()
print(" Metrics:")
print(f" {'Metric':<14s} {'Value':>12s} {'Change':>8s}")
print(f" {'-'*14} {'-'*12} {'-'*8}")
print(f" {'HCVA':<14s} {fmt_dollar(hcva):>12s} {fmt_pct(hcva, p_hcva):>8s}")
print(f" {'HCROI':<14s} {hcroi:>12.2f}x {fmt_pct(hcroi, p_hcroi):>8s}")
print(f" {'QIPS':<14s} {qips:>12.4f} {fmt_pct(qips, p_qips):>8s}")
print(f" {'Empl. Value':<14s} {ev:>12.4f}")
print(f" {'Headcount':<14s} {headcount:>12d} {fmt_pct(headcount, p_hc):>8s}")
print(f" {'Engagement':<14s} {engagement:>12.1f} {fmt_pct(engagement, p_eng):>8s}")
if reward is not None:
print(f"\n Quarterly Reward: {reward:+.4f}")
def print_final_summary(
total_steps: int,
final_score: float,
metric_history: List[dict],
quarterly_rewards: List[float],
all_events: List[str],
) -> None:
print()
print("=" * 64)
print(" FINAL EPISODE SUMMARY")
print("=" * 64)
print(f" Total Steps: {total_steps}")
print(f" Final Score: {final_score:.4f}")
if len(metric_history) >= 2:
first = metric_history[0]
last = metric_history[-1]
print(f"\n Metric Trajectory (Baseline -> Q6):")
h0, h1 = first.get("hcva", 0), last.get("hcva", 0)
r0, r1 = first.get("hcroi", 0), last.get("hcroi", 0)
q0, q1 = first.get("qips", {}).get("composite", 0), last.get("qips", {}).get("composite", 0)
print(f" HCVA: {fmt_dollar(h0):>10s} -> {fmt_dollar(h1):>10s} ({fmt_pct(h1, h0)})")
print(f" HCROI: {r0:>10.2f}x -> {r1:>10.2f}x ({fmt_pct(r1, r0)})")
print(f" QIPS: {q0:>10.4f} -> {q1:>10.4f} ({fmt_pct(q1, q0)})")
if quarterly_rewards:
formatted = ", ".join(f"{r:+.4f}" for r in quarterly_rewards)
print(f"\n Quarterly Rewards: [{formatted}]")
if all_events:
print(f"\n Events Log ({len(all_events)} total):")
for ev in all_events:
print(f" - {ev}")
print("=" * 64)
print()
# ── Phase Execution ─────────────────────────────────────────────────
def run_scanning(env: HRProductivityEnvironment) -> dict:
"""Execute scanning phase: query all departments, employees, metrics, financials."""
print_phase("scanning")
dept_data = {}
# Query each department
for dept_name in DEPARTMENTS:
obs = env.step(HRAction(action_type="query_department", department=dept_name))
d = obs.data or {}
dept_data[dept_name] = d
print(
f" [SCAN] {dept_name:12s}: {d.get('headcount', '?'):>3} employees, "
f"perf {d.get('avg_performance', 0):.1f}, "
f"engage {d.get('avg_engagement', 0):.0f}, "
f"risk {d.get('avg_flight_risk', 0):.2f}"
)
# Query high performers for promotion candidates
obs = env.step(HRAction(
action_type="query_employees", parameters={"min_performance": 4.0}
))
high_performers = (obs.data or {}).get("employees", [])
promote_ids = [
e["id"] for e in high_performers if e.get("level", 5) < 5
][:3]
# Query low performers for termination candidates
obs = env.step(HRAction(
action_type="query_employees", parameters={"min_flight_risk": 0.0}
))
all_emps = (obs.data or {}).get("employees", [])
terminate_ids = [
e["id"] for e in all_emps if e.get("performance_score", 5.0) < 1.8
][:2]
# Calculate all metrics
obs = env.step(HRAction(action_type="calculate_metric", metric_name="all"))
all_metrics = obs.data or {}
hcva = all_metrics.get("hcva", 0)
hcroi = all_metrics.get("hcroi", 0)
qips_c = all_metrics.get("qips", {}).get("composite", 0)
print(f" [METRIC] HCVA: {fmt_dollar(hcva)} | HCROI: {hcroi:.2f}x | QIPS: {qips_c:.4f}")
# Review financials
obs = env.step(HRAction(action_type="review_financials"))
fin = obs.data or {}
hr_budget = fin.get("hr_budget_remaining", 0)
print(f" [FINANCE] Revenue: {fmt_dollar(fin.get('revenue', 0))} | "
f"Profit: {fmt_dollar(fin.get('profit', 0))} | "
f"HR Budget: {fmt_dollar(hr_budget)}")
# Advance to planning
env.step(HRAction(action_type="advance_phase"))
return {
"depts": dept_data,
"metrics": all_metrics,
"financials": fin,
"hr_budget": hr_budget,
"promote_ids": promote_ids,
"terminate_ids": terminate_ids,
}
def plan_quarter(
env: HRProductivityEnvironment,
scan_data: dict,
quarter: int,
prev_events: List[str],
) -> dict:
"""Execute planning phase with data-driven decisions."""
print_phase("planning")
depts = scan_data["depts"]
hr_budget = scan_data["hr_budget"]
plan = {}
# Rank departments by flight risk (descending) and performance (ascending)
dept_list = [
{
"name": name,
"headcount": d.get("headcount", 0),
"avg_performance": d.get("avg_performance", 3.0),
"avg_engagement": d.get("avg_engagement", 70),
"avg_flight_risk": d.get("avg_flight_risk", 0.1),
}
for name, d in depts.items()
]
by_risk = sorted(dept_list, key=lambda d: d["avg_flight_risk"], reverse=True)
by_perf = sorted(dept_list, key=lambda d: d["avg_performance"])
by_headcount = sorted(
[d for d in dept_list if d["name"] != "HR"],
key=lambda d: d["headcount"],
)
# 1. Hiring: spread across the 2 smallest non-HR depts to offset turnover
hiring_dept = by_headcount[0]["name"]
# React to competitor poaching events
for ev in prev_events:
if "poaching" in ev.lower() or "competitor" in ev.lower():
for d in dept_list:
if d["name"].lower() in ev.lower():
hiring_dept = d["name"]
break
# Hire aggressively to offset ~20% quarterly turnover
total_hc = sum(d["headcount"] for d in dept_list)
hire_count = max(3, min(15, int(total_hc * 0.12)))
obs = env.step(HRAction(
action_type="set_hiring_target", department=hiring_dept, count=hire_count
))
plan["hiring"] = {"dept": hiring_dept, "count": hire_count}
print(f" [PLAN] Hiring target: {hiring_dept} +{hire_count}")
# Also set hiring target for second-smallest dept
hiring_dept2 = by_headcount[1]["name"] if len(by_headcount) > 1 else hiring_dept
hire_count2 = max(2, min(10, int(total_hc * 0.08)))
if hiring_dept2 != hiring_dept:
obs = env.step(HRAction(
action_type="set_hiring_target", department=hiring_dept2, count=hire_count2
))
plan["hiring2"] = {"dept": hiring_dept2, "count": hire_count2}
print(f" [PLAN] Hiring target: {hiring_dept2} +{hire_count2}")
# 2. Training: lowest-performing department
training_dept = by_perf[0]["name"]
training_amount = min(hr_budget * 0.30, 100_000)
if training_amount > 0:
obs = env.step(HRAction(
action_type="set_training_budget", department=training_dept, amount=training_amount
))
# Update budget tracking from observation
if obs.data and "hr_budget_remaining" in obs.data:
hr_budget = obs.data["hr_budget_remaining"]
else:
hr_budget -= training_amount
plan["training"] = {"dept": training_dept, "amount": training_amount}
print(f" [PLAN] Training budget: {fmt_dollar(training_amount)} -> {training_dept}")
# 3. Compensation: highest flight-risk department, +3%
comp_dept = by_risk[0]["name"]
comp_pct = 3.0
obs = env.step(HRAction(
action_type="set_compensation_policy", department=comp_dept, amount=comp_pct
))
plan["compensation"] = {"dept": comp_dept, "pct": comp_pct}
print(f" [PLAN] Compensation: +{comp_pct:.0f}% for {comp_dept}")
# 4. Retention: highest flight-risk dept (if budget allows)
retention_dept = by_risk[0]["name"]
retention_amount = min(hr_budget * 0.20, 50_000)
if retention_amount > 1000:
obs = env.step(HRAction(
action_type="set_retention_program", department=retention_dept, amount=retention_amount
))
plan["retention"] = {"dept": retention_dept, "amount": retention_amount}
print(f" [PLAN] Retention program: {retention_dept} ({fmt_dollar(retention_amount)})")
# Advance to producing
env.step(HRAction(action_type="advance_phase"))
return plan
def execute_quarter(
env: HRProductivityEnvironment,
plan: dict,
scan_data: dict,
) -> None:
"""Execute producing phase: hire, train, promote, optionally terminate."""
print_phase("producing")
# 1. Execute hiring (primary dept)
hiring = plan.get("hiring", {})
if hiring:
obs = env.step(HRAction(
action_type="execute_hiring",
department=hiring["dept"],
count=hiring["count"],
))
cost = (obs.data or {}).get("cost", 0)
print(f" [EXEC] Hired {hiring['count']} in {hiring['dept']} (cost: {fmt_dollar(cost)})")
# 1b. Execute hiring (secondary dept)
hiring2 = plan.get("hiring2", {})
if hiring2:
obs = env.step(HRAction(
action_type="execute_hiring",
department=hiring2["dept"],
count=hiring2["count"],
))
cost = (obs.data or {}).get("cost", 0)
print(f" [EXEC] Hired {hiring2['count']} in {hiring2['dept']} (cost: {fmt_dollar(cost)})")
# 2. Execute training
training = plan.get("training", {})
if training:
obs = env.step(HRAction(
action_type="execute_training",
department=training["dept"],
amount=20, # 20 hours per employee
))
cost = (obs.data or {}).get("cost", 0)
hc = (obs.data or {}).get("headcount", "?")
print(f" [EXEC] Training: {training['dept']}, 20 hrs x {hc} employees (cost: {fmt_dollar(cost)})")
# 3. Promote top performers (identified during scanning)
promote_ids = scan_data.get("promote_ids", [])
if promote_ids:
obs = env.step(HRAction(action_type="execute_promotion", employee_ids=promote_ids))
promoted = sum(
1 for p in (obs.data or {}).get("promotions", []) if p.get("success")
)
print(f" [EXEC] Promoted {promoted} high-performing employees")
else:
print(" [EXEC] No promotion candidates found")
# 4. Terminate underperformers (identified during scanning)
terminate_ids = scan_data.get("terminate_ids", [])
if terminate_ids:
obs = env.step(HRAction(action_type="execute_termination", employee_ids=terminate_ids))
terminated = sum(
1 for t in (obs.data or {}).get("terminations", []) if t.get("success")
)
print(f" [EXEC] Terminated {terminated} underperformers")
else:
print(" [EXEC] No underperformers to terminate")
# Advance to controlling
env.step(HRAction(action_type="advance_phase"))
def run_controlling(env: HRProductivityEnvironment, q: int) -> dict:
"""Execute controlling phase: submit report and advance quarter."""
print_phase("controlling")
# Submit report
env.step(HRAction(action_type="submit_report"))
print(f" [REPORT] Q{q} report submitted")
# Advance quarter
obs = env.step(HRAction(action_type="advance_quarter"))
return {
"obs": obs,
"data": obs.data or {},
"reward": obs.reward,
"done": obs.done,
}
# ── Main Demo Runner ────────────────────────────────────────────────
def run_demo(seed: int = 42, size: int = 300, scenario: Optional[str] = None) -> float:
"""Run a full 6-quarter demo episode with a heuristic strategy."""
print_header(seed, size, scenario)
env = HRProductivityEnvironment()
kwargs: Dict[str, Any] = {"size": size}
if scenario:
kwargs["scenario"] = scenario
obs = env.reset(seed=seed, **kwargs)
print_baseline(obs.data or {})
# Track state across quarters
baseline_metrics = (obs.data or {}).get("baseline_metrics", {})
metric_history = [baseline_metrics]
quarterly_rewards: List[float] = []
all_events: List[str] = []
prev_events: List[str] = []
total_steps = 0
for q in range(1, 7):
print_quarter_header(q)
# Scanning
scan_data = run_scanning(env)
# Planning
plan = plan_quarter(env, scan_data, q, prev_events)
# Producing
execute_quarter(env, plan, scan_data)
# Controlling
result = run_controlling(env, q)
data = result["data"]
# Collect results
reward = result["reward"]
events = data.get("events", [])
turnover = data.get("turnover", {})
departed = turnover.get("departed_count", 0)
metrics = data.get("metrics") or data.get("last_quarter_metrics", {})
all_events.extend(events)
prev_events = events
if reward is not None and not result["done"]:
quarterly_rewards.append(reward)
if metrics:
metric_history.append(metrics)
# Print quarter results
if events:
print(f"\n Events:")
for ev in events:
print(f" - {ev}")
if departed:
print(f" Turnover: {departed} employees departed")
prev_m = metric_history[-2] if len(metric_history) >= 2 else None
if metrics:
print_metrics_table(metrics, prev_m, reward)
total_steps = data.get("total_steps", total_steps)
if result["done"]:
final_score = data.get("final_score", reward or 0.0)
total_steps = data.get("total_steps", total_steps)
# Collect all quarterly rewards from final data
qr = data.get("quarterly_rewards", quarterly_rewards)
if qr:
quarterly_rewards = qr
print_final_summary(
total_steps, final_score, metric_history, quarterly_rewards, all_events
)
return final_score
# Should not reach here, but just in case
return 0.0
def _pct_change(curr: float, prev: float) -> Optional[float]:
if prev == 0:
return None
return round((curr - prev) / abs(prev) * 100, 2)
def run_demo_json(
seed: int = 42, size: int = 300, scenario: Optional[str] = None
) -> Dict[str, Any]:
"""Run a full 6-quarter demo and return structured JSON results.
This is the API-friendly version of run_demo() — no print output,
returns all data as a dict suitable for JSON serialization.
"""
env = HRProductivityEnvironment()
kwargs: Dict[str, Any] = {"size": size}
if scenario:
kwargs["scenario"] = scenario
obs = env.reset(seed=seed, **kwargs)
obs_data = obs.data or {}
baseline_metrics = obs_data.get("baseline_metrics", {})
baseline_summary = obs_data.get("company_summary", {})
metric_history = [baseline_metrics]
quarterly_rewards: List[float] = []
all_events: List[str] = []
prev_events: List[str] = []
quarters: List[Dict[str, Any]] = []
for q in range(1, 7):
scan_data = run_scanning.__wrapped__(env) if hasattr(run_scanning, '__wrapped__') else _scan_quiet(env)
plan = _plan_quiet(env, scan_data, q, prev_events)
_exec_quiet(env, plan, scan_data)
result = _control_quiet(env)
data = result["data"]
reward = result["reward"]
events = data.get("events", [])
turnover = data.get("turnover", {})
departed = turnover.get("departed_count", 0)
metrics = data.get("metrics") or data.get("last_quarter_metrics", {})
all_events.extend(events)
prev_events = events
prev_m = metric_history[-1] if metric_history else {}
quarter_result: Dict[str, Any] = {
"quarter": q,
"actions": {
"hiring": plan.get("hiring"),
"hiring2": plan.get("hiring2"),
"training": plan.get("training"),
"compensation": plan.get("compensation"),
"retention": plan.get("retention"),
},
"events": events,
"turnover_count": departed,
"quarterly_reward": reward,
}
if metrics:
metric_history.append(metrics)
quarter_result["metrics"] = {
"hcva": metrics.get("hcva", 0),
"hcroi": metrics.get("hcroi", 0),
"qips": metrics.get("qips", {}).get("composite", 0),
"employee_value": metrics.get("employee_value", 0),
"headcount": metrics.get("snapshot", {}).get("headcount", 0),
"avg_engagement": round(metrics.get("snapshot", {}).get("avg_engagement", 0), 1),
}
quarter_result["changes"] = {
"hcva_pct": _pct_change(metrics.get("hcva", 0), prev_m.get("hcva", 0)),
"hcroi_pct": _pct_change(metrics.get("hcroi", 0), prev_m.get("hcroi", 0)),
"qips_pct": _pct_change(
metrics.get("qips", {}).get("composite", 0),
prev_m.get("qips", {}).get("composite", 0),
),
}
if reward is not None and not result["done"]:
quarterly_rewards.append(reward)
quarters.append(quarter_result)
if result["done"]:
final_score = data.get("final_score", reward or 0.0)
qr = data.get("quarterly_rewards", quarterly_rewards)
if qr:
quarterly_rewards = qr
break
first_m = metric_history[0] if metric_history else {}
last_m = metric_history[-1] if metric_history else {}
return {
"seed": seed,
"size": size,
"scenario": scenario,
"baseline": {
"headcount": baseline_summary.get("total_headcount", 0),
"revenue": baseline_summary.get("revenue", 0),
"profit": baseline_summary.get("profit", 0),
"hcva": first_m.get("hcva", 0),
"hcroi": first_m.get("hcroi", 0),
"qips": first_m.get("qips", {}).get("composite", 0),
},
"final": {
"score": final_score,
"total_steps": data.get("total_steps", 0),
"hcva": last_m.get("hcva", 0),
"hcroi": last_m.get("hcroi", 0),
"qips": last_m.get("qips", {}).get("composite", 0),
"hcva_change_pct": _pct_change(last_m.get("hcva", 0), first_m.get("hcva", 0)),
"hcroi_change_pct": _pct_change(last_m.get("hcroi", 0), first_m.get("hcroi", 0)),
"qips_change_pct": _pct_change(
last_m.get("qips", {}).get("composite", 0),
first_m.get("qips", {}).get("composite", 0),
),
},
"quarterly_rewards": quarterly_rewards,
"events": all_events,
"quarters": quarters,
}
def _scan_quiet(env: HRProductivityEnvironment) -> dict:
"""Scanning phase — no print output."""
dept_data = {}
for dept_name in DEPARTMENTS:
obs = env.step(HRAction(action_type="query_department", department=dept_name))
dept_data[dept_name] = obs.data or {}
obs = env.step(HRAction(
action_type="query_employees", parameters={"min_performance": 4.0}
))
high_performers = (obs.data or {}).get("employees", [])
promote_ids = [e["id"] for e in high_performers if e.get("level", 5) < 5][:3]
obs = env.step(HRAction(
action_type="query_employees", parameters={"min_flight_risk": 0.0}
))
all_emps = (obs.data or {}).get("employees", [])
terminate_ids = [
e["id"] for e in all_emps if e.get("performance_score", 5.0) < 1.8
][:2]
obs = env.step(HRAction(action_type="calculate_metric", metric_name="all"))
all_metrics = obs.data or {}
obs = env.step(HRAction(action_type="review_financials"))
fin = obs.data or {}
env.step(HRAction(action_type="advance_phase"))
return {
"depts": dept_data,
"metrics": all_metrics,
"financials": fin,
"hr_budget": fin.get("hr_budget_remaining", 0),
"promote_ids": promote_ids,
"terminate_ids": terminate_ids,
}
def _plan_quiet(
env: HRProductivityEnvironment, scan_data: dict, quarter: int, prev_events: List[str]
) -> dict:
"""Planning phase — no print output."""
depts = scan_data["depts"]
hr_budget = scan_data["hr_budget"]
plan: Dict[str, Any] = {}
dept_list = [
{
"name": name,
"headcount": d.get("headcount", 0),
"avg_performance": d.get("avg_performance", 3.0),
"avg_engagement": d.get("avg_engagement", 70),
"avg_flight_risk": d.get("avg_flight_risk", 0.1),
}
for name, d in depts.items()
]
by_risk = sorted(dept_list, key=lambda d: d["avg_flight_risk"], reverse=True)
by_perf = sorted(dept_list, key=lambda d: d["avg_performance"])
by_headcount = sorted(
[d for d in dept_list if d["name"] != "HR"], key=lambda d: d["headcount"]
)
hiring_dept = by_headcount[0]["name"]
for ev in prev_events:
if "poaching" in ev.lower() or "competitor" in ev.lower():
for d in dept_list:
if d["name"].lower() in ev.lower():
hiring_dept = d["name"]
break
total_hc = sum(d["headcount"] for d in dept_list)
hire_count = max(3, min(15, int(total_hc * 0.12)))
env.step(HRAction(action_type="set_hiring_target", department=hiring_dept, count=hire_count))
plan["hiring"] = {"dept": hiring_dept, "count": hire_count}
hiring_dept2 = by_headcount[1]["name"] if len(by_headcount) > 1 else hiring_dept
hire_count2 = max(2, min(10, int(total_hc * 0.08)))
if hiring_dept2 != hiring_dept:
env.step(HRAction(action_type="set_hiring_target", department=hiring_dept2, count=hire_count2))
plan["hiring2"] = {"dept": hiring_dept2, "count": hire_count2}
training_dept = by_perf[0]["name"]
training_amount = min(hr_budget * 0.30, 100_000)
if training_amount > 0:
obs = env.step(HRAction(
action_type="set_training_budget", department=training_dept, amount=training_amount
))
if obs.data and "hr_budget_remaining" in obs.data:
hr_budget = obs.data["hr_budget_remaining"]
else:
hr_budget -= training_amount
plan["training"] = {"dept": training_dept, "amount": training_amount}
comp_dept = by_risk[0]["name"]
env.step(HRAction(action_type="set_compensation_policy", department=comp_dept, amount=3.0))
plan["compensation"] = {"dept": comp_dept, "pct": 3.0}
retention_dept = by_risk[0]["name"]
retention_amount = min(hr_budget * 0.20, 50_000)
if retention_amount > 1000:
env.step(HRAction(
action_type="set_retention_program", department=retention_dept, amount=retention_amount
))
plan["retention"] = {"dept": retention_dept, "amount": retention_amount}
env.step(HRAction(action_type="advance_phase"))
return plan
def _exec_quiet(
env: HRProductivityEnvironment, plan: dict, scan_data: dict
) -> None:
"""Producing phase — no print output."""
hiring = plan.get("hiring", {})
if hiring:
env.step(HRAction(action_type="execute_hiring", department=hiring["dept"], count=hiring["count"]))
hiring2 = plan.get("hiring2", {})
if hiring2:
env.step(HRAction(action_type="execute_hiring", department=hiring2["dept"], count=hiring2["count"]))
training = plan.get("training", {})
if training:
env.step(HRAction(action_type="execute_training", department=training["dept"], amount=20))
promote_ids = scan_data.get("promote_ids", [])
if promote_ids:
env.step(HRAction(action_type="execute_promotion", employee_ids=promote_ids))
terminate_ids = scan_data.get("terminate_ids", [])
if terminate_ids:
env.step(HRAction(action_type="execute_termination", employee_ids=terminate_ids))
env.step(HRAction(action_type="advance_phase"))
def _control_quiet(env: HRProductivityEnvironment) -> dict:
"""Controlling phase — no print output."""
env.step(HRAction(action_type="submit_report"))
obs = env.step(HRAction(action_type="advance_quarter"))
return {
"obs": obs,
"data": obs.data or {},
"reward": obs.reward,
"done": obs.done,
}
def main() -> None:
parser = argparse.ArgumentParser(
description="Run a full 6-quarter HR simulation demo with heuristic strategy"
)
parser.add_argument("--seed", type=int, default=42, help="Random seed (default: 42)")
parser.add_argument("--size", type=int, default=300, help="Company size (default: 300)")
parser.add_argument(
"--scenario",
choices=["high_eng_turnover", "budget_cuts", "rapid_growth", "balanced_optimization"],
help="Scenario variant to run",
)
parser.add_argument(
"--all-scenarios", action="store_true",
help="Run all 4 task scenarios and compare results",
)
args = parser.parse_args()
if args.all_scenarios:
print()
print("=" * 64)
print(" Running all 4 scenarios...")
print("=" * 64)
results = []
for scn in SCENARIOS:
score = run_demo(seed=scn["seed"], size=scn["size"], scenario=scn["name"])
results.append((scn["name"], score))
print()
print("=" * 64)
print(" SCENARIO COMPARISON")
print("=" * 64)
print(f" {'Scenario':<28s} {'Score':>8s}")
print(f" {'-'*28} {'-'*8}")
for name, score in results:
print(f" {name:<28s} {score:>8.4f}")
avg = sum(s for _, s in results) / len(results)
print(f" {'-'*28} {'-'*8}")
print(f" {'Average':<28s} {avg:>8.4f}")
print("=" * 64)
else:
score = run_demo(seed=args.seed, size=args.size, scenario=args.scenario)
print(f" Exit score: {score:.4f}")
if __name__ == "__main__":
main()