anupamagarwal001/amc_allocator_env
0
1# Copyright (c) Meta Platforms, Inc. and affiliates.2# All rights reserved.3#4# This source code is licensed under the BSD-style license found in the5# LICENSE file in the root directory of this source tree.6 7"""Data models for the Round 2 committee environment."""8 9from __future__ import annotations10 11import math12from typing import Dict, List, Literal13 14from openenv.core.env_server.types import Action, Observation, State15from pydantic import Field, field_validator16 17 18ActionType = Literal[19 "query_research",20 "query_risk",21 "allocate",22 "revise_allocation",23 "hold",24 "move_to_cash",25]26 27 28class PortfolioAction(Action):29 """Committee-style action emitted by the Portfolio Manager."""30 31 action_type: ActionType = Field(32 default="hold",33 description="Type of committee action to take this step.",34 )35 allocation_template: str | None = Field(36 default=None,37 description="Optional discrete allocation template name.",38 )39 query_target: str | None = Field(40 default=None,41 description="Optional research focus such as an asset ticker or SECTOR.",42 )43 target_weights: Dict[str, float] = Field(44 default_factory=dict,45 description="Optional custom asset weights. Cash is implied by any unallocated weight.",46 )47 reason: str | None = Field(48 default=None,49 description="Optional explanation for the committee decision.",50 )51 52 @field_validator("target_weights")53 @classmethod54 def validate_target_weights(cls, value: Dict[str, float]) -> Dict[str, float]:55 cleaned: Dict[str, float] = {}56 for asset, weight in value.items():57 if not asset or not asset.strip():58 raise ValueError("asset names must be non-empty")59 if not math.isfinite(weight):60 raise ValueError(f"weight for {asset!r} must be finite")61 cleaned[asset.strip().upper()] = float(weight)62 return cleaned63 64 @field_validator("query_target")65 @classmethod66 def normalize_query_target(cls, value: str | None) -> str | None:67 if value is None:68 return None69 normalized = value.strip().upper()70 return normalized or None71 72 @field_validator("allocation_template")73 @classmethod74 def normalize_allocation_template(cls, value: str | None) -> str | None:75 if value is None:76 return None77 normalized = value.strip().lower()78 return normalized or None79 80 81class AllocatorObservation(Observation):82 """Observation visible to the trainable Portfolio Manager."""83 84 task_id: str = Field(..., description="Active task identifier.")85 task_description: str = Field(..., description="Human-readable task summary.")86 step_index: int = Field(..., ge=0, description="Current decision step index.")87 steps_remaining: int = Field(..., ge=0, description="Remaining decisions in the episode.")88 prices: Dict[str, float] = Field(..., description="Current asset prices.")89 signals: Dict[str, float] = Field(90 ...,91 description="Current market signals in the range [-1.0, 1.0].",92 )93 current_weights: Dict[str, float] = Field(94 default_factory=dict,95 description="Current invested asset weights before the next action.",96 )97 cash_weight: float = Field(..., ge=0.0, le=1.0, description="Current cash allocation.")98 portfolio_value: float = Field(..., ge=0.0, description="Current portfolio NAV.")99 turnover: float = Field(..., ge=0.0, description="Turnover generated by the last action.")100 risk_metrics: Dict[str, float] = Field(101 default_factory=dict,102 description="Latest market and risk summary metrics.",103 )104 active_constraints: Dict[str, float] = Field(105 default_factory=dict,106 description="Current mandate and risk constraints.",107 )108 reward_components: Dict[str, float] = Field(109 default_factory=dict,110 description="Last step reward component breakdown.",111 )112 research_notes: List[str] = Field(113 default_factory=list,114 description="Recent research notes available to the PM.",115 )116 risk_notes: List[str] = Field(117 default_factory=list,118 description="Recent risk notes available to the PM.",119 )120 available_allocation_templates: Dict[str, str] = Field(121 default_factory=dict,122 description="Discrete allocation templates exposed by the environment.",123 )124 outstanding_risk_alert: bool = Field(125 default=False,126 description="Whether the PM is currently under elevated risk pressure.",127 )128 queries_remaining: int = Field(129 default=0,130 ge=0,131 description="Remaining advisory queries in the episode.",132 )133 134 135class AllocatorState(State):136 """Extended environment state exposed through the OpenEnv state endpoint."""137 138 task_id: str = Field(default="guided_allocation", description="Active task identifier.")139 task_description: str = Field(default="", description="Human-readable task summary.")140 total_steps: int = Field(default=0, ge=0, description="Total decisions in the episode.")141 current_step: int = Field(default=0, ge=0, description="Current completed step count.")142 assets: List[str] = Field(default_factory=list, description="Tradable asset universe.")143 holdings: Dict[str, float] = Field(144 default_factory=dict,145 description="Current position sizes in asset units.",146 )147 current_weights: Dict[str, float] = Field(148 default_factory=dict,149 description="Current invested asset weights.",150 )151 cash_weight: float = Field(default=1.0, ge=0.0, le=1.0, description="Current cash weight.")152 portfolio_value: float = Field(default=1.0, ge=0.0, description="Current portfolio NAV.")153 nav_history: List[float] = Field(default_factory=list, description="Episode NAV history.")154 turnover_history: List[float] = Field(155 default_factory=list,156 description="Per-step turnover values.",157 )158 reward_history: List[float] = Field(159 default_factory=list,160 description="Per-step shaped rewards.",161 )162 portfolio_return_history: List[float] = Field(163 default_factory=list,164 description="Per-step realized portfolio returns before costs.",165 )166 signal_alignment_history: List[float] = Field(167 default_factory=list,168 description="Per-step alignment between allocated weights and market signals.",169 )170 information_usage_history: List[float] = Field(171 default_factory=list,172 description="Per-step usage score for analyst information.",173 )174 risk_response_history: List[float] = Field(175 default_factory=list,176 description="Per-step response quality to risk pressure.",177 )178 compliance_penalty_history: List[float] = Field(179 default_factory=list,180 description="Per-step compliance penalty values.",181 )182 reward_component_history: List[Dict[str, float]] = Field(183 default_factory=list,184 description="Reward breakdown history.",185 )186 risk_metrics: Dict[str, float] = Field(187 default_factory=dict,188 description="Latest risk summary emitted in the observation.",189 )190 active_constraints: Dict[str, float] = Field(191 default_factory=dict,192 description="Current risk and mandate constraints.",193 )194 research_notes_history: List[str] = Field(195 default_factory=list,196 description="Research commentary visible to the PM.",197 )198 risk_notes_history: List[str] = Field(199 default_factory=list,200 description="Risk commentary visible to the PM.",201 )202 action_history: List[str] = Field(203 default_factory=list,204 description="Portfolio Manager action history.",205 )206 allocation_template_history: List[str] = Field(207 default_factory=list,208 description="Chosen allocation templates over time.",209 )210 hidden_regime_history: List[str] = Field(211 default_factory=list,212 description="Internal regime labels used by the environment.",213 )214 latest_research_view: Dict[str, float] = Field(215 default_factory=dict,216 description="Latest noisy research score snapshot.",217 )218 latest_risk_view: Dict[str, float] = Field(219 default_factory=dict,220 description="Latest risk guidance snapshot.",221 )222 query_count: int = Field(default=0, ge=0, description="Number of advisory queries used.")223 outstanding_risk_alert: bool = Field(224 default=False,225 description="Whether the PM currently faces elevated risk pressure.",226 )227 compliance_violations: int = Field(228 default=0,229 ge=0,230 description="Count of compliance or mandate breaches.",231 )232 last_turnover: float = Field(default=0.0, ge=0.0, description="Latest turnover.")233 last_reward: float = Field(default=0.0, description="Latest shaped reward.")234 last_action_type: str = Field(default="reset", description="Latest action type.")235 last_action_reason: str | None = Field(236 default=None,237 description="Optional latest action rationale.",238 )239 