CoolFace
Apppublic

anupamagarwal001/amc_allocator_env

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