Spaces:
Paused
Paused
Sync Space with HCM-21-Private main (72f5f30): reward recalibration, seeded RNG streams, MCP integration, poster artifacts
670ccf0 verified | """Core HR Productivity Environment β OpenEnv-compliant reset/step/state.""" | |
| from __future__ import annotations | |
| import random | |
| import uuid | |
| from typing import Any, Dict, List, Optional | |
| from openenv.core.env_server.interfaces import Environment | |
| from hr_env.models import HRAction, HRObservation, HRState | |
| from hr_env.server.company import Company, DEPARTMENT_NAMES, RECRUITING_COST_PER_HIRE | |
| from hr_env.server.data_gen import generate_company | |
| from hr_env.server.employee import Employee | |
| from hr_env.server.events import apply_events | |
| from hr_env.server.metrics import ( | |
| compute_all_metrics, | |
| compute_employee_value, | |
| compute_hcva, | |
| compute_hcroi, | |
| compute_qips, | |
| ) | |
| from hr_env.server.phases import PhaseManager | |
| from hr_env.server.scoring import compute_final_score, compute_quarterly_reward | |
| MAX_QUARTERS = 6 | |
| MAX_STEPS = 300 # Safety limit | |
| class HRProductivityEnvironment(Environment[HRAction, HRObservation, HRState]): | |
| """6-quarter HR simulation where an LLM agent acts as Chief Human Capital Officer.""" | |
| def __init__(self, **kwargs: Any): | |
| super().__init__(**kwargs) | |
| self._company: Optional[Company] = None | |
| self._phase_manager = PhaseManager() | |
| self._rng = random.Random(42) | |
| self._episode_id: Optional[str] = None | |
| self._step_count = 0 | |
| self._quarter_step_count = 0 | |
| self._metric_history: List[Dict[str, Any]] = [] | |
| self._previous_snapshot: Optional[Dict] = None | |
| self._event_log: List[str] = [] | |
| self._active_initiatives: List[Dict[str, Any]] = [] | |
| self._quarterly_rewards: List[float] = [] | |
| self._done = False | |
| def reset( | |
| self, seed: Optional[int] = None, episode_id: Optional[str] = None, **kwargs: Any | |
| ) -> HRObservation: | |
| """Reset the environment for a new episode.""" | |
| if seed is None: | |
| seed = random.randint(0, 2**31) | |
| self._rng = random.Random(seed) | |
| self._episode_id = episode_id or str(uuid.uuid4())[:12] | |
| self._step_count = 0 | |
| self._quarter_step_count = 0 | |
| self._metric_history = [] | |
| self._previous_snapshot = None | |
| self._event_log = [] | |
| self._active_initiatives = [] | |
| self._quarterly_rewards = [] | |
| self._done = False | |
| self._phase_manager = PhaseManager() | |
| # Generate company from kwargs or defaults | |
| company_size = kwargs.get("size", 300) | |
| scenario = kwargs.get("scenario", None) | |
| self._company = generate_company( | |
| seed=seed, | |
| size=company_size, | |
| name=kwargs.get("company_name", "Simulated Corp"), | |
| ) | |
| # Apply scenario modifications | |
| if scenario: | |
| self._apply_scenario(scenario) | |
| # Compute baseline metrics | |
| baseline = compute_all_metrics(self._company) | |
| baseline["profit"] = self._company.compute_profit() | |
| self._metric_history.append(baseline) | |
| self._previous_snapshot = baseline.get("snapshot") | |
| return HRObservation( | |
| done=False, | |
| reward=None, | |
| message=( | |
| f"Welcome, Chief Human Capital Officer of {self._company.name}. " | |
| f"You manage {self._company.total_headcount} employees across " | |
| f"{len(self._company.departments)} departments. " | |
| f"Your goal: maximize human capital value over 6 quarters. " | |
| f"Current quarter: Q1. Phase: Scanning. " | |
| f"Begin by gathering data about your organization." | |
| ), | |
| data={ | |
| "company_summary": self._company.summary(), | |
| "baseline_metrics": baseline, | |
| }, | |
| available_actions=self._phase_manager.get_available_actions(), | |
| current_phase="scanning", | |
| current_quarter=1, | |
| step_in_quarter=0, | |
| warnings=[], | |
| ) | |
| def step( | |
| self, action: HRAction, timeout_s: Optional[float] = None, **kwargs: Any | |
| ) -> HRObservation: | |
| """Execute one action and return the observation.""" | |
| if self._done: | |
| return self._make_done_observation("Episode is already complete.") | |
| if self._company is None: | |
| return HRObservation( | |
| done=False, | |
| reward=None, | |
| message="Error: Environment not initialized. Call reset() first.", | |
| available_actions=[], | |
| current_phase="scanning", | |
| current_quarter=1, | |
| step_in_quarter=0, | |
| warnings=["Environment not initialized"], | |
| ) | |
| self._step_count += 1 | |
| self._quarter_step_count += 1 | |
| # Safety limit | |
| if self._step_count >= MAX_STEPS: | |
| return self._finalize_episode("Maximum step limit reached.") | |
| # Validate action is allowed in current phase | |
| action_type = action.action_type | |
| if not self._phase_manager.is_valid_action(action_type): | |
| available = self._phase_manager.get_available_actions() | |
| return HRObservation( | |
| done=False, | |
| reward=None, | |
| message=( | |
| f"Invalid action '{action_type}' in {self._phase_manager.current_phase} phase. " | |
| f"Available actions: {', '.join(available)}" | |
| ), | |
| available_actions=available, | |
| current_phase=self._phase_manager.current_phase, | |
| current_quarter=self._company.quarter + 1, | |
| step_in_quarter=self._quarter_step_count, | |
| warnings=[f"Invalid action: {action_type}"], | |
| ) | |
| # Record the action | |
| self._phase_manager.record_action(action_type) | |
| # Dispatch to handler | |
| handler = self._get_handler(action_type) | |
| return handler(action) | |
| def state(self) -> HRState: | |
| """Return the full environment state.""" | |
| if self._company is None: | |
| return HRState() | |
| return HRState( | |
| episode_id=self._episode_id, | |
| step_count=self._step_count, | |
| current_quarter=self._company.quarter + 1, | |
| current_phase=self._phase_manager.current_phase, | |
| total_steps=self._step_count, | |
| company_snapshot=self._company.snapshot(), | |
| metric_history=self._metric_history, | |
| budget_remaining={ | |
| "hr_budget": self._company.hr_budget_remaining, | |
| **{ | |
| d.name: d.training_budget - d.training_budget_spent | |
| for d in self._company.departments.values() | |
| }, | |
| }, | |
| active_initiatives=self._active_initiatives, | |
| recent_events=self._event_log[-5:], | |
| ) | |
| # ββ Action Handlers βββββββββββββββββββββββββββββββββββββββββββββ | |
| def _get_handler(self, action_type: str): | |
| handlers = { | |
| # Scanning | |
| "query_department": self._handle_query_department, | |
| "query_employees": self._handle_query_employees, | |
| "calculate_metric": self._handle_calculate_metric, | |
| "review_financials": self._handle_review_financials, | |
| # Planning | |
| "set_hiring_target": self._handle_set_hiring_target, | |
| "set_training_budget": self._handle_set_training_budget, | |
| "set_compensation_policy": self._handle_set_compensation_policy, | |
| "set_retention_program": self._handle_set_retention_program, | |
| # Producing | |
| "execute_hiring": self._handle_execute_hiring, | |
| "execute_promotion": self._handle_execute_promotion, | |
| "execute_transfer": self._handle_execute_transfer, | |
| "execute_training": self._handle_execute_training, | |
| "execute_termination": self._handle_execute_termination, | |
| # Control | |
| "advance_phase": self._handle_advance_phase, | |
| "advance_quarter": self._handle_advance_quarter, | |
| "submit_report": self._handle_submit_report, | |
| } | |
| return handlers[action_type] | |
| def _handle_query_department(self, action: HRAction) -> HRObservation: | |
| dept_name = action.department | |
| if not dept_name or dept_name not in self._company.departments: | |
| return self._obs(f"Department '{dept_name}' not found. Available: {', '.join(DEPARTMENT_NAMES)}") | |
| dept = self._company.departments[dept_name] | |
| return self._obs( | |
| f"Department: {dept_name}", | |
| data=dept.to_dict(), | |
| ) | |
| def _handle_query_employees(self, action: HRAction) -> HRObservation: | |
| filters = action.parameters or {} | |
| dept_name = action.department | |
| employees = self._company.all_active_employees | |
| if dept_name: | |
| dept = self._company.get_department(dept_name) | |
| if not dept: | |
| return self._obs(f"Department '{dept_name}' not found.") | |
| employees = dept.active_employees | |
| # Apply filters | |
| if "min_performance" in filters: | |
| employees = [e for e in employees if e.performance_score >= filters["min_performance"]] | |
| if "max_flight_risk" in filters: | |
| employees = [e for e in employees if e.flight_risk <= filters["max_flight_risk"]] | |
| if "min_flight_risk" in filters: | |
| employees = [e for e in employees if e.flight_risk >= filters["min_flight_risk"]] | |
| if "level" in filters: | |
| employees = [e for e in employees if e.level == filters["level"]] | |
| # Limit to 20 for observation size | |
| total = len(employees) | |
| employees = employees[:20] | |
| return self._obs( | |
| f"Found {total} employees (showing {len(employees)})", | |
| data={ | |
| "total": total, | |
| "employees": [e.summary() for e in employees], | |
| }, | |
| ) | |
| def _handle_calculate_metric(self, action: HRAction) -> HRObservation: | |
| metric = action.metric_name or "all" | |
| if metric == "hcva": | |
| value = compute_hcva(self._company) | |
| return self._obs(f"HCVA: ${value:,.2f} per FTE", data={"hcva": round(value, 2)}) | |
| elif metric == "hcroi": | |
| value = compute_hcroi(self._company) | |
| return self._obs(f"HCROI: {value:.4f}x", data={"hcroi": round(value, 4)}) | |
| elif metric == "qips": | |
| value = compute_qips(self._company) | |
| return self._obs(f"QIPS composite: {value['composite']:.4f}", data={"qips": value}) | |
| elif metric == "employee_value": | |
| value = compute_employee_value(self._company) | |
| return self._obs(f"Employee Value: {value:.4f}", data={"employee_value": value}) | |
| elif metric == "five_indexes": | |
| current_snapshot = { | |
| "employment_cost": self._company.total_employment_cost, | |
| "time_to_fill": 30, | |
| "headcount": self._company.total_headcount, | |
| "avg_performance": sum(e.performance_score for e in self._company.all_active_employees) / max(1, len(self._company.all_active_employees)), | |
| "avg_engagement": sum(e.engagement for e in self._company.all_active_employees) / max(1, len(self._company.all_active_employees)), | |
| } | |
| from hr_env.server.metrics import compute_five_indexes | |
| value = compute_five_indexes(current_snapshot, self._previous_snapshot) | |
| return self._obs("Five Indexes of Change", data={"five_indexes": value}) | |
| else: | |
| all_metrics = compute_all_metrics(self._company, self._previous_snapshot) | |
| return self._obs("All metrics computed", data=all_metrics) | |
| def _handle_review_financials(self, action: HRAction) -> HRObservation: | |
| revenue = self._company.compute_revenue() | |
| profit = self._company.compute_profit() | |
| return self._obs( | |
| f"Revenue: ${revenue:,.0f} | Profit: ${profit:,.0f} | Employment Cost: ${self._company.total_employment_cost:,.0f}", | |
| data={ | |
| "revenue": round(revenue, 2), | |
| "profit": round(profit, 2), | |
| "total_employment_cost": round(self._company.total_employment_cost, 2), | |
| "total_salary_cost": round(self._company.total_salary_cost, 2), | |
| "hr_budget_remaining": round(self._company.hr_budget_remaining, 2), | |
| "departments": { | |
| name: {"headcount": d.headcount, "employment_cost": round(d.employment_cost(), 2)} | |
| for name, d in self._company.departments.items() | |
| }, | |
| }, | |
| ) | |
| def _handle_set_hiring_target(self, action: HRAction) -> HRObservation: | |
| dept_name = action.department | |
| count = action.count or 0 | |
| if not dept_name or dept_name not in self._company.departments: | |
| return self._obs(f"Invalid department. Available: {', '.join(DEPARTMENT_NAMES)}") | |
| dept = self._company.departments[dept_name] | |
| dept.hiring_target = count | |
| cost_estimate = count * RECRUITING_COST_PER_HIRE | |
| return self._obs( | |
| f"Hiring target set: {count} new hires for {dept_name} (est. cost: ${cost_estimate:,.0f})", | |
| data={"department": dept_name, "hiring_target": count, "estimated_cost": cost_estimate}, | |
| ) | |
| def _handle_set_training_budget(self, action: HRAction) -> HRObservation: | |
| dept_name = action.department | |
| amount = action.amount or 0 | |
| if not dept_name or dept_name not in self._company.departments: | |
| return self._obs(f"Invalid department. Available: {', '.join(DEPARTMENT_NAMES)}") | |
| if amount > self._company.hr_budget_remaining: | |
| return self._obs( | |
| f"Insufficient HR budget. Requested: ${amount:,.0f}, Available: ${self._company.hr_budget_remaining:,.0f}", | |
| warnings=["Budget exceeded"], | |
| ) | |
| dept = self._company.departments[dept_name] | |
| dept.training_budget = amount | |
| self._company.hr_budget_remaining -= amount | |
| return self._obs( | |
| f"Training budget set: ${amount:,.0f} for {dept_name}. HR budget remaining: ${self._company.hr_budget_remaining:,.0f}", | |
| data={"department": dept_name, "training_budget": amount, "hr_budget_remaining": self._company.hr_budget_remaining}, | |
| ) | |
| def _handle_set_compensation_policy(self, action: HRAction) -> HRObservation: | |
| dept_name = action.department | |
| adjustment = action.amount or 0 # Percentage | |
| if not dept_name or dept_name not in self._company.departments: | |
| return self._obs(f"Invalid department. Available: {', '.join(DEPARTMENT_NAMES)}") | |
| dept = self._company.departments[dept_name] | |
| dept.compensation_adjustment = adjustment | |
| estimated_cost = dept.total_salary * (adjustment / 100) | |
| return self._obs( | |
| f"Compensation adjustment set: {adjustment:+.1f}% for {dept_name} (est. cost: ${estimated_cost:,.0f}/yr)", | |
| data={"department": dept_name, "adjustment_pct": adjustment, "estimated_annual_cost": round(estimated_cost, 2)}, | |
| ) | |
| def _handle_set_retention_program(self, action: HRAction) -> HRObservation: | |
| dept_name = action.department | |
| budget = action.amount or 0 | |
| if not dept_name or dept_name not in self._company.departments: | |
| return self._obs(f"Invalid department. Available: {', '.join(DEPARTMENT_NAMES)}") | |
| if budget > self._company.hr_budget_remaining: | |
| return self._obs( | |
| f"Insufficient HR budget. Requested: ${budget:,.0f}, Available: ${self._company.hr_budget_remaining:,.0f}", | |
| warnings=["Budget exceeded"], | |
| ) | |
| dept = self._company.departments[dept_name] | |
| dept.retention_program_active = True | |
| dept.retention_budget = budget | |
| self._company.hr_budget_remaining -= budget | |
| self._active_initiatives.append({ | |
| "type": "retention_program", | |
| "department": dept_name, | |
| "budget": budget, | |
| "quarter_started": self._company.quarter + 1, | |
| }) | |
| return self._obs( | |
| f"Retention program activated for {dept_name} with ${budget:,.0f} budget.", | |
| data={"department": dept_name, "retention_budget": budget}, | |
| ) | |
| def _handle_execute_hiring(self, action: HRAction) -> HRObservation: | |
| dept_name = action.department | |
| count = action.count or 0 | |
| if not dept_name or dept_name not in self._company.departments: | |
| return self._obs(f"Invalid department. Available: {', '.join(DEPARTMENT_NAMES)}") | |
| dept = self._company.departments[dept_name] | |
| if count <= 0: | |
| return self._obs("Must hire at least 1 employee.") | |
| # Cost check | |
| hire_cost = count * RECRUITING_COST_PER_HIRE | |
| if hire_cost > self._company.hr_budget_remaining: | |
| return self._obs( | |
| f"Insufficient budget for {count} hires. Cost: ${hire_cost:,.0f}, Available: ${self._company.hr_budget_remaining:,.0f}", | |
| warnings=["Budget exceeded"], | |
| ) | |
| from hr_env.server.data_gen import ROLES, SALARY_BASE, DEPT_SKILLS | |
| import numpy as np | |
| from faker import Faker | |
| # Derive both RNGs from the seeded episode RNG so hires are reproducible. | |
| hire_seed = self._rng.randint(0, 2**31) | |
| Faker.seed(hire_seed) | |
| fake = Faker() | |
| rng_np = np.random.default_rng(hire_seed) | |
| hired = [] | |
| for _ in range(count): | |
| level = 1 # New hires start at level 1-2 | |
| if self._rng.random() < 0.3: | |
| level = 2 | |
| roles = ROLES.get(dept_name, [("Employee", 1)]) | |
| role_name = next((r for r, l in roles if l == level), roles[0][0]) | |
| base_salary = SALARY_BASE.get(level, 55000) | |
| salary = round(float(rng_np.lognormal(np.log(base_salary), 0.12)), -2) | |
| skills_pool = DEPT_SKILLS.get(dept_name, ["general"]) | |
| n_skills = min(len(skills_pool), int(rng_np.integers(2, 4))) | |
| emp = Employee( | |
| id=f"{dept_name[:3].lower()}_{self._rng.randint(5000, 9999)}", | |
| name=fake.name(), | |
| department=dept_name, | |
| role=role_name, | |
| level=level, | |
| salary=salary, | |
| tenure_months=0, | |
| performance_score=float(np.clip(rng_np.normal(3.0, 0.5), 1.0, 5.0)), | |
| engagement=float(rng_np.beta(8, 2) * 100), # New hires tend to be enthusiastic | |
| skills=list(rng_np.choice(skills_pool, size=n_skills, replace=False)), | |
| ) | |
| emp.update_flight_risk() | |
| dept.employees.append(emp) | |
| hired.append(emp.summary()) | |
| self._company.hr_budget_remaining -= hire_cost | |
| return self._obs( | |
| f"Hired {count} employees in {dept_name}. Cost: ${hire_cost:,.0f}.", | |
| data={"hired": hired, "cost": hire_cost, "department": dept_name}, | |
| ) | |
| def _handle_execute_promotion(self, action: HRAction) -> HRObservation: | |
| if not action.employee_ids: | |
| return self._obs("No employee IDs provided for promotion.") | |
| results = [] | |
| for emp_id in action.employee_ids: | |
| emp = self._company.get_employee(emp_id) | |
| if not emp or not emp.is_active: | |
| results.append({"id": emp_id, "success": False, "reason": "Employee not found or inactive"}) | |
| continue | |
| cost = emp.salary * 0.15 # Promotion raises cost 15% | |
| success = emp.promote() | |
| results.append({ | |
| "id": emp_id, "name": emp.name, "success": success, | |
| "new_level": emp.level, "new_salary": round(emp.salary, 2), | |
| }) | |
| return self._obs( | |
| f"Promotion results: {sum(1 for r in results if r.get('success'))}/{len(results)} successful.", | |
| data={"promotions": results}, | |
| ) | |
| def _handle_execute_transfer(self, action: HRAction) -> HRObservation: | |
| if not action.employee_ids or not action.department: | |
| return self._obs("Provide employee_ids and target department for transfer.") | |
| target_dept = action.department | |
| if target_dept not in self._company.departments: | |
| return self._obs(f"Target department '{target_dept}' not found.") | |
| results = [] | |
| for emp_id in action.employee_ids: | |
| emp = self._company.get_employee(emp_id) | |
| if not emp or not emp.is_active: | |
| results.append({"id": emp_id, "success": False, "reason": "Not found"}) | |
| continue | |
| source_dept = self._company.find_employee_department(emp_id) | |
| if source_dept == target_dept: | |
| results.append({"id": emp_id, "success": False, "reason": "Already in target dept"}) | |
| continue | |
| # Move employee | |
| self._company.departments[source_dept].employees.remove(emp) | |
| emp.department = target_dept | |
| self._company.departments[target_dept].employees.append(emp) | |
| emp.update_engagement(-5) # Transfer stress | |
| results.append({"id": emp_id, "name": emp.name, "success": True, "from": source_dept, "to": target_dept}) | |
| return self._obs( | |
| f"Transfer results: {sum(1 for r in results if r.get('success'))}/{len(results)} completed.", | |
| data={"transfers": results}, | |
| ) | |
| def _handle_execute_training(self, action: HRAction) -> HRObservation: | |
| dept_name = action.department | |
| if not dept_name or dept_name not in self._company.departments: | |
| return self._obs(f"Invalid department. Available: {', '.join(DEPARTMENT_NAMES)}") | |
| dept = self._company.departments[dept_name] | |
| hours = action.amount or 20 # Default 20 hours | |
| cost_per_hour = 50 # $50/hr/employee | |
| total_cost = hours * cost_per_hour * dept.headcount | |
| if total_cost > dept.training_budget - dept.training_budget_spent: | |
| remaining = dept.training_budget - dept.training_budget_spent | |
| return self._obs( | |
| f"Training cost ${total_cost:,.0f} exceeds remaining budget ${remaining:,.0f} for {dept_name}.", | |
| warnings=["Training budget exceeded"], | |
| ) | |
| dept.training_budget_spent += total_cost | |
| for emp in dept.active_employees: | |
| emp.apply_training(hours) | |
| return self._obs( | |
| f"Training: {dept.headcount} employees in {dept_name} received {hours:.0f} hours. Cost: ${total_cost:,.0f}.", | |
| data={"department": dept_name, "hours": hours, "headcount": dept.headcount, "cost": total_cost}, | |
| ) | |
| def _handle_execute_termination(self, action: HRAction) -> HRObservation: | |
| if not action.employee_ids: | |
| return self._obs("No employee IDs provided for termination.") | |
| results = [] | |
| for emp_id in action.employee_ids: | |
| emp = self._company.get_employee(emp_id) | |
| if not emp or not emp.is_active: | |
| results.append({"id": emp_id, "success": False, "reason": "Not found or already inactive"}) | |
| continue | |
| emp.is_active = False | |
| # Termination impacts department morale | |
| dept = self._company.get_department(emp.department) | |
| if dept: | |
| for colleague in dept.active_employees: | |
| colleague.update_engagement(-3) | |
| results.append({"id": emp_id, "name": emp.name, "success": True, "department": emp.department}) | |
| return self._obs( | |
| f"Terminated {sum(1 for r in results if r.get('success'))}/{len(results)} employees.", | |
| data={"terminations": results}, | |
| ) | |
| def _handle_advance_phase(self, action: HRAction) -> HRObservation: | |
| try: | |
| new_phase = self._phase_manager.advance_phase() | |
| return self._obs( | |
| f"Advanced to {new_phase} phase.", | |
| data={"phase": new_phase, "phase_summary": self._phase_manager.get_phase_summary()}, | |
| ) | |
| except ValueError as e: | |
| return self._obs(str(e), warnings=[str(e)]) | |
| def _handle_advance_quarter(self, action: HRAction) -> HRObservation: | |
| if self._phase_manager.current_phase != "controlling": | |
| return self._obs( | |
| "Cannot advance quarter: must be in controlling phase. Current: " + self._phase_manager.current_phase, | |
| warnings=["Not in controlling phase"], | |
| ) | |
| min_req = 1 | |
| if self._phase_manager.actions_in_phase < min_req: | |
| return self._obs( | |
| f"Complete at least {min_req} controlling action before advancing quarter.", | |
| warnings=["Minimum controlling actions not met"], | |
| ) | |
| # Apply compensation adjustments | |
| for dept in self._company.departments.values(): | |
| if dept.compensation_adjustment != 0: | |
| for emp in dept.active_employees: | |
| emp.salary *= (1 + dept.compensation_adjustment / 100) | |
| if dept.compensation_adjustment > 0: | |
| emp.update_engagement(2) | |
| emp.flight_risk = max(0.02, emp.flight_risk - 0.03) | |
| # Apply retention programs | |
| for dept in self._company.departments.values(): | |
| if dept.retention_program_active: | |
| for emp in dept.active_employees: | |
| emp.update_engagement(5) | |
| emp.flight_risk = max(0.02, emp.flight_risk - 0.10) | |
| # Advance company (handles turnover, resets) | |
| turnover_info = self._company.advance_quarter() | |
| # Apply stochastic events | |
| event_descriptions = apply_events(self._company, self._rng) | |
| self._event_log.extend(event_descriptions) | |
| # Compute metrics | |
| metrics = compute_all_metrics(self._company, self._previous_snapshot) | |
| metrics["profit"] = self._company.compute_profit() | |
| metrics["turnover"] = turnover_info | |
| # Quarterly reward | |
| if len(self._metric_history) > 0: | |
| reward = compute_quarterly_reward(metrics, self._metric_history[-1]) | |
| else: | |
| reward = 0.0 | |
| self._quarterly_rewards.append(reward) | |
| self._metric_history.append(metrics) | |
| self._previous_snapshot = metrics.get("snapshot") | |
| # Check if episode is done (Q6 completed) | |
| quarter_num = self._company.quarter | |
| if quarter_num >= MAX_QUARTERS: | |
| return self._finalize_episode( | |
| f"Q{quarter_num} complete. All 6 quarters finished!", | |
| quarterly_reward=reward, | |
| turnover_info=turnover_info, | |
| events=event_descriptions, | |
| metrics=metrics, | |
| ) | |
| # Reset phase for new quarter | |
| self._phase_manager.reset_for_new_quarter() | |
| self._quarter_step_count = 0 | |
| warnings = [] | |
| if len(turnover_info["departed"]) > 0: | |
| warnings.append(f"{len(turnover_info['departed'])} employees departed this quarter.") | |
| if event_descriptions: | |
| warnings.extend(event_descriptions) | |
| return HRObservation( | |
| done=False, | |
| reward=reward, | |
| message=( | |
| f"Q{quarter_num} complete. Advancing to Q{quarter_num + 1}. " | |
| f"Turnover: {len(turnover_info['departed'])} departed. " | |
| f"{'Events: ' + '; '.join(event_descriptions) if event_descriptions else 'No special events.'}" | |
| ), | |
| data={ | |
| "quarter_completed": quarter_num, | |
| "quarterly_reward": reward, | |
| "metrics": metrics, | |
| "turnover": {"departed_count": len(turnover_info["departed"])}, | |
| "events": event_descriptions, | |
| "next_quarter": quarter_num + 1, | |
| }, | |
| available_actions=self._phase_manager.get_available_actions(), | |
| current_phase="scanning", | |
| current_quarter=quarter_num + 1, | |
| step_in_quarter=0, | |
| warnings=warnings, | |
| ) | |
| def _handle_submit_report(self, action: HRAction) -> HRObservation: | |
| metrics = compute_all_metrics(self._company, self._previous_snapshot) | |
| company_snapshot = self._company.snapshot() | |
| return self._obs( | |
| f"Q{self._company.quarter + 1} HR Report submitted.", | |
| data={ | |
| "report": { | |
| "quarter": self._company.quarter + 1, | |
| "total_steps": self._step_count, | |
| "company_snapshot": company_snapshot, | |
| "metrics": metrics, | |
| "active_initiatives": self._active_initiatives, | |
| "recent_events": self._event_log[-3:], | |
| }, | |
| }, | |
| ) | |
| # ββ Helpers ββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def _obs(self, message: str, data: Optional[dict] = None, warnings: Optional[list] = None) -> HRObservation: | |
| """Create a standard observation.""" | |
| return HRObservation( | |
| done=False, | |
| reward=None, | |
| message=message, | |
| data=data, | |
| available_actions=self._phase_manager.get_available_actions(), | |
| current_phase=self._phase_manager.current_phase, | |
| current_quarter=self._company.quarter + 1 if self._company else 1, | |
| step_in_quarter=self._quarter_step_count, | |
| warnings=warnings or [], | |
| ) | |
| def _make_done_observation(self, message: str) -> HRObservation: | |
| final_score = compute_final_score(self._metric_history) | |
| return HRObservation( | |
| done=True, | |
| reward=final_score, | |
| message=message + f" Final score: {final_score:.4f}", | |
| data={"final_score": final_score, "metric_history": self._metric_history}, | |
| available_actions=[], | |
| current_phase="done", | |
| current_quarter=MAX_QUARTERS, | |
| step_in_quarter=0, | |
| warnings=[], | |
| ) | |
| def _finalize_episode( | |
| self, | |
| message: str, | |
| quarterly_reward: float = 0.0, | |
| turnover_info: Optional[dict] = None, | |
| events: Optional[list] = None, | |
| metrics: Optional[dict] = None, | |
| ) -> HRObservation: | |
| """Finalize the episode and compute final score.""" | |
| self._done = True | |
| final_score = compute_final_score(self._metric_history) | |
| return HRObservation( | |
| done=True, | |
| reward=final_score, | |
| message=f"{message} Final episode score: {final_score:.4f}", | |
| data={ | |
| "final_score": final_score, | |
| "quarterly_rewards": self._quarterly_rewards, | |
| "total_steps": self._step_count, | |
| "metric_history": self._metric_history, | |
| "last_quarter_metrics": metrics, | |
| "last_quarter_turnover": turnover_info, | |
| "events": events or [], | |
| }, | |
| available_actions=[], | |
| current_phase="done", | |
| current_quarter=MAX_QUARTERS, | |
| step_in_quarter=0, | |
| warnings=[], | |
| ) | |
| def _apply_scenario(self, scenario: str) -> None: | |
| """Apply scenario-specific modifications to the generated company.""" | |
| if scenario == "high_eng_turnover": | |
| eng = self._company.departments.get("Engineering") | |
| if eng: | |
| for emp in eng.active_employees: | |
| emp.flight_risk = min(0.95, emp.flight_risk + 0.20) | |
| emp.update_engagement(-10) | |
| elif scenario == "budget_cuts": | |
| self._company.hr_budget *= 0.5 | |
| self._company.hr_budget_remaining = self._company.hr_budget / 4 | |
| elif scenario == "rapid_growth": | |
| self._company.base_revenue *= 1.3 | |
| self._company.market_modifier = 1.15 | |
| # More hiring needed | |
| for dept in self._company.departments.values(): | |
| # Remove some employees to create understaffing | |
| to_remove = max(1, int(len(dept.active_employees) * 0.15)) | |
| for emp in dept.active_employees[:to_remove]: | |
| emp.is_active = False | |
| elif scenario == "balanced_optimization": | |
| # Default company with slightly below-average metrics | |
| for emp in self._company.all_active_employees: | |
| emp.performance_score = max(1.0, emp.performance_score - 0.3) | |
| emp.engagement = max(20, emp.engagement - 8) | |