CoolFace
Apppublic

anupamagarwal001/amc_allocator_env

sourceHugging Faceupdated 5mo agoView on Hugging Face
0likes
tasks.py238 linesDownload Raw Back to root
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