anupamagarwal001/amc_allocator_env
0
1"""Round 2 task definitions and scenario loading."""2 3from __future__ import annotations4 5from dataclasses import dataclass6from typing import Dict7 8try:9 from .data import MARKET_SCENARIOS10except ImportError: # pragma: no cover11 from data import MARKET_SCENARIOS12 13DEFAULT_TASK_ID = "guided_allocation"14TASK_ORDER = (15 "guided_allocation",16 "research_risk_conflict",17 "regime_shift_recovery",18 "mandate_drift",19)20 21 22@dataclass(frozen=True)23class TaskScenario:24 """Deterministic committee-style scenario configuration."""25 26 task_id: str27 display_name: str28 summary: str29 sector: str30 assets: list[str]31 steps: int32 transaction_cost_bps: float33 risk_aversion: float34 drawdown_penalty: float35 cash_bias: float36 query_cost: float37 max_queries: int38 research_noise: float39 risk_alert_threshold: float40 base_constraints: Dict[str, float]41 constraint_schedule: Dict[int, Dict[str, float]]42 prices: list[Dict[str, float]]43 signals: list[Dict[str, float]]44 regimes: list[str]45 46 def __post_init__(self) -> None:47 if self.steps <= 0:48 raise ValueError(f"{self.task_id} must define at least one step")49 if len(self.prices) != self.steps + 1:50 raise ValueError(51 f"{self.task_id} requires {self.steps + 1} price rows, got {len(self.prices)}"52 )53 if len(self.signals) != self.steps:54 raise ValueError(55 f"{self.task_id} requires {self.steps} signal rows, got {len(self.signals)}"56 )57 if len(self.regimes) != self.steps:58 raise ValueError(59 f"{self.task_id} requires {self.steps} regime labels, got {len(self.regimes)}"60 )61 62 63def _clone_rows(rows: list[Dict[str, float]]) -> list[Dict[str, float]]:64 return [{asset: float(value) for asset, value in row.items()} for row in rows]65 66 67def _regimes(steps: int, labels: list[tuple[int, str]]) -> list[str]:68 output: list[str] = []69 cursor = 070 for stop, label in labels:71 bounded_stop = min(stop, steps)72 while cursor < bounded_stop:73 output.append(label)74 cursor += 175 while cursor < steps:76 output.append(labels[-1][1])77 cursor += 178 return output79 80 81def load_task_scenarios() -> Dict[str, TaskScenario]:82 """Build the Round 2 task set from the offline market dataset."""83 84 sector = str(MARKET_SCENARIOS["sector"])85 assets = [str(asset) for asset in MARKET_SCENARIOS["assets"]]86 raw_tasks = MARKET_SCENARIOS["tasks"]87 88 signal_following = raw_tasks["signal_following"]89 noisy_market = raw_tasks["noisy_market"]90 regime_shift = raw_tasks["regime_shift"]91 92 scenarios = {93 "guided_allocation": TaskScenario(94 task_id="guided_allocation",95 display_name="Guided Allocation",96 summary=(97 "The Portfolio Manager learns the committee workflow in a stable market "98 "with clean research input and light risk pressure."99 ),100 sector=sector,101 assets=assets,102 steps=int(signal_following["steps"]),103 transaction_cost_bps=float(signal_following["transaction_cost_bps"]),104 risk_aversion=float(signal_following["risk_aversion"]),105 drawdown_penalty=float(signal_following["drawdown_penalty"]),106 cash_bias=float(signal_following["cash_bias"]),107 query_cost=0.0012,108 max_queries=8,109 research_noise=0.12,110 risk_alert_threshold=0.52,111 base_constraints={112 "max_single_asset_weight": 0.50,113 "min_cash_weight": 0.05,114 "max_turnover": 0.65,115 },116 constraint_schedule={},117 prices=_clone_rows(signal_following["prices"]),118 signals=_clone_rows(signal_following["signals"]),119 regimes=_regimes(int(signal_following["steps"]), [(30, "steady")]),120 ),121 "research_risk_conflict": TaskScenario(122 task_id="research_risk_conflict",123 display_name="Research vs Risk Conflict",124 summary=(125 "Bullish analyst views collide with tighter risk constraints and "126 "fragile market internals."127 ),128 sector=sector,129 assets=assets,130 steps=int(noisy_market["steps"]),131 transaction_cost_bps=float(noisy_market["transaction_cost_bps"]),132 risk_aversion=float(noisy_market["risk_aversion"]),133 drawdown_penalty=float(noisy_market["drawdown_penalty"]),134 cash_bias=float(noisy_market["cash_bias"]),135 query_cost=0.0018,136 max_queries=10,137 research_noise=0.22,138 risk_alert_threshold=0.46,139 base_constraints={140 "max_single_asset_weight": 0.38,141 "min_cash_weight": 0.08,142 "max_turnover": 0.45,143 },144 constraint_schedule={145 12: {"max_single_asset_weight": 0.32},146 28: {"min_cash_weight": 0.14},147 },148 prices=_clone_rows(noisy_market["prices"]),149 signals=_clone_rows(noisy_market["signals"]),150 regimes=_regimes(151 int(noisy_market["steps"]),152 [(15, "crowded"), (30, "fragile"), (45, "fragile")],153 ),154 ),155 "regime_shift_recovery": TaskScenario(156 task_id="regime_shift_recovery",157 display_name="Regime Shift Recovery",158 summary=(159 "The committee must respond to hidden regime deterioration, de-risk, "160 "and then selectively re-risk as conditions improve."161 ),162 sector=sector,163 assets=assets,164 steps=int(regime_shift["steps"]),165 transaction_cost_bps=float(regime_shift["transaction_cost_bps"]),166 risk_aversion=float(regime_shift["risk_aversion"]),167 drawdown_penalty=float(regime_shift["drawdown_penalty"]),168 cash_bias=float(regime_shift["cash_bias"]),169 query_cost=0.0016,170 max_queries=10,171 research_noise=0.18,172 risk_alert_threshold=0.42,173 base_constraints={174 "max_single_asset_weight": 0.42,175 "min_cash_weight": 0.10,176 "max_turnover": 0.40,177 },178 constraint_schedule={179 20: {"min_cash_weight": 0.18},180 40: {"min_cash_weight": 0.08},181 },182 prices=_clone_rows(regime_shift["prices"]),183 signals=_clone_rows(regime_shift["signals"]),184 regimes=_regimes(185 int(regime_shift["steps"]),186 [(20, "expansion"), (40, "stress"), (60, "repair")],187 ),188 ),189 "mandate_drift": TaskScenario(190 task_id="mandate_drift",191 display_name="Mandate Drift",192 summary=(193 "The PM must manage a long-horizon allocation process while compliance "194 "rules and cash mandates tighten mid-episode."195 ),196 sector=sector,197 assets=assets,198 steps=int(noisy_market["steps"]),199 transaction_cost_bps=float(noisy_market["transaction_cost_bps"]) + 2.0,200 risk_aversion=float(noisy_market["risk_aversion"]) + 0.08,201 drawdown_penalty=float(noisy_market["drawdown_penalty"]) + 0.05,202 cash_bias=float(noisy_market["cash_bias"]) + 0.03,203 query_cost=0.0015,204 max_queries=9,205 research_noise=0.18,206 risk_alert_threshold=0.44,207 base_constraints={208 "max_single_asset_weight": 0.45,209 "min_cash_weight": 0.05,210 "max_turnover": 0.50,211 },212 constraint_schedule={213 15: {"max_single_asset_weight": 0.28, "min_cash_weight": 0.12},214 30: {"max_single_asset_weight": 0.35, "min_cash_weight": 0.18, "max_turnover": 0.35},215 },216 prices=_clone_rows(noisy_market["prices"]),217 signals=_clone_rows(noisy_market["signals"]),218 regimes=_regimes(219 int(noisy_market["steps"]),220 [(15, "baseline"), (30, "tightening"), (45, "oversight")],221 ),222 ),223 }224 return scenarios225 226 227TASK_SCENARIOS = load_task_scenarios()228 229 230def get_task_scenario(task_id: str) -> TaskScenario:231 """Return the requested task scenario or raise on unknown task ids."""232 233 normalized_task_id = task_id.strip().lower()234 if normalized_task_id not in TASK_SCENARIOS:235 available = ", ".join(TASK_ORDER)236 raise KeyError(f"Unknown task_id={task_id!r}. Expected one of: {available}")237 return TASK_SCENARIOS[normalized_task_id]238 