Spaces:
Sleeping
Sleeping
Commit ·
2a048ac
1
Parent(s): 1b5bc6d
Added logic for 1_Day_move and 30_Day_Move and created directory for get_figures task
Browse files- .gitignore +2 -0
- codewiki.md +704 -0
- earnings_analyst/models.py +5 -0
- evaluate.py +11 -0
- models.py +9 -0
- scratch/inspect_dataset.py +26 -0
- tasks/1_day_move/grader.py +7 -5
- tasks/1_day_move/spec.py +36 -6
- tasks/30_day_move/grader.py +7 -5
- tasks/30_day_move/spec.py +36 -6
- tasks/get_figures/__init__.py +4 -0
- tasks/get_figures/grader.py +14 -0
- tasks/get_figures/spec.py +29 -0
- tasks/grading.py +52 -0
- tasks/registry.py +3 -1
.gitignore
CHANGED
|
@@ -26,3 +26,5 @@ openenv_earnings_analyst.egg-info/
|
|
| 26 |
# Huggingface spaces doesn't allow pdfs anywhere in git history
|
| 27 |
*.pdf
|
| 28 |
*.parquet
|
|
|
|
|
|
|
|
|
| 26 |
# Huggingface spaces doesn't allow pdfs anywhere in git history
|
| 27 |
*.pdf
|
| 28 |
*.parquet
|
| 29 |
+
|
| 30 |
+
results/
|
codewiki.md
ADDED
|
@@ -0,0 +1,704 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Earnings Lens — Code Wiki
|
| 2 |
+
|
| 3 |
+
> Package: `openenv-earnings_analyst` v0.1.0
|
| 4 |
+
> Repo: `RudrakshNanavaty/earnings-lens`
|
| 5 |
+
> Python ≥ 3.12 · Dependency manager: `uv`
|
| 6 |
+
|
| 7 |
+
---
|
| 8 |
+
|
| 9 |
+
## Table of Contents
|
| 10 |
+
|
| 11 |
+
1. [Project Overview](#1-project-overview)
|
| 12 |
+
2. [Repository Layout](#2-repository-layout)
|
| 13 |
+
3. [Architecture & Data Flow](#3-architecture--data-flow)
|
| 14 |
+
4. [Module Reference](#4-module-reference)
|
| 15 |
+
- [Root Package (`earnings_analyst`)](#41-root-package-earnings_analyst)
|
| 16 |
+
- [`models.py`](#42-modelspy)
|
| 17 |
+
- [`client.py`](#43-clientpy)
|
| 18 |
+
- [`environment_config.py`](#44-environment_configpy)
|
| 19 |
+
- [`server/`](#45-server)
|
| 20 |
+
- [`tasks/`](#46-tasks)
|
| 21 |
+
5. [Task System Deep-Dive](#5-task-system-deep-dive)
|
| 22 |
+
- [`TaskSpec` schema](#51-taskspec-schema)
|
| 23 |
+
- [Task registry](#52-task-registry)
|
| 24 |
+
- [Grading helpers](#53-grading-helpers)
|
| 25 |
+
- [Task inventory](#54-task-inventory)
|
| 26 |
+
6. [Environment Lifecycle](#6-environment-lifecycle)
|
| 27 |
+
7. [Inference & Evaluation Scripts](#7-inference--evaluation-scripts)
|
| 28 |
+
8. [Configuration Reference](#8-configuration-reference)
|
| 29 |
+
9. [Deployment](#9-deployment)
|
| 30 |
+
10. [Adding a New Task](#10-adding-a-new-task)
|
| 31 |
+
11. [Key Design Decisions & Gotchas](#11-key-design-decisions--gotchas)
|
| 32 |
+
|
| 33 |
+
---
|
| 34 |
+
|
| 35 |
+
## 1. Project Overview
|
| 36 |
+
|
| 37 |
+
**Earnings Lens** is an [OpenEnv](https://github.com/meta-pytorch/OpenEnv) environment that wraps a Hugging Face earnings-call dataset as an RL interaction loop. An LLM agent receives one row of earnings-call data (transcripts, press releases, market numerics) and must predict a task-specific target (e.g. sentiment label, or a price move). The server scores the prediction and returns a reward.
|
| 38 |
+
|
| 39 |
+
**Core concepts:**
|
| 40 |
+
|
| 41 |
+
| Term | Meaning |
|
| 42 |
+
|------|---------|
|
| 43 |
+
| **Episode** | One dataset row. `reset()` samples it; `step(prediction)` scores it. Always terminal after one step (`done=True`). |
|
| 44 |
+
| **Task** | Named configuration (`task_id`) with a `TaskSpec` (what columns to expose, what to predict) and a `grade()` function. |
|
| 45 |
+
| **Active task** | Chosen at server startup via `EARNINGS_ANALYST_TASK_ID`. Cannot be changed per-connection. |
|
| 46 |
+
| **Implementation gate** | `spec["implemented"] = False` blocks `reset()` with `TaskNotImplementedError`, preventing accidental evaluation of stub tasks. |
|
| 47 |
+
|
| 48 |
+
---
|
| 49 |
+
|
| 50 |
+
## 2. Repository Layout
|
| 51 |
+
|
| 52 |
+
```
|
| 53 |
+
earnings-lens/
|
| 54 |
+
├── __init__.py # Public API: EarningsAnalystEnv, Action, Observation
|
| 55 |
+
├── models.py # Pydantic Action + Observation models
|
| 56 |
+
├── client.py # WebSocket EnvClient subclass
|
| 57 |
+
├── environment_config.py # Dataset ID/file; re-exports DEFAULT_TASK, TASKS
|
| 58 |
+
├── inference.py # Example: single-episode OpenAI inference
|
| 59 |
+
├── evaluate.py # Batch evaluation over N episodes
|
| 60 |
+
├── main.py # Placeholder ("Hello from earnings-analyst!")
|
| 61 |
+
├── openenv.yaml # OpenEnv Spaces manifest (app: server.app:app, port 8000)
|
| 62 |
+
├── pyproject.toml # Package metadata, deps, console script `server`
|
| 63 |
+
├── uv.lock # Locked deps for `uv sync`
|
| 64 |
+
├── .env.example # Template for local .env
|
| 65 |
+
│
|
| 66 |
+
├── server/
|
| 67 |
+
│ ├── __init__.py # Re-exports EarningsAnalystEnvironment
|
| 68 |
+
│ ├── app.py # FastAPI app factory + main() entrypoint
|
| 69 |
+
│ ├── earnings_analyst_environment.py # Core Environment: reset / step
|
| 70 |
+
│ ├── dataset_loader.py # Module-level HF dataset singleton
|
| 71 |
+
│ ├── Dockerfile # Production Docker image
|
| 72 |
+
│ └── requirements.txt # Optional pins for Docker / tooling
|
| 73 |
+
│
|
| 74 |
+
└── tasks/
|
| 75 |
+
├── __init__.py # Re-exports registry symbols + TaskNotImplementedError
|
| 76 |
+
├── types.py # TaskSpec TypedDict
|
| 77 |
+
├── exceptions.py # TaskNotImplementedError
|
| 78 |
+
├── grading.py # Shared helpers: grade_ordinal, grade_exact
|
| 79 |
+
├── loader.py # load_task_subpackage() for digit-prefixed folders
|
| 80 |
+
├── registry.py # Central task registry: TASKS, GRADERS, get_grader()
|
| 81 |
+
│
|
| 82 |
+
├── sentiment_label/ # ✅ Implemented
|
| 83 |
+
│ ├── spec.py
|
| 84 |
+
│ ├── grader.py
|
| 85 |
+
│ └── __init__.py
|
| 86 |
+
├── 1_day_move/ # 🚧 Stub
|
| 87 |
+
│ ├── spec.py
|
| 88 |
+
│ ├── grader.py
|
| 89 |
+
│ └── __init__.py
|
| 90 |
+
├── 30_day_move/ # 🚧 Stub
|
| 91 |
+
│ ├── spec.py
|
| 92 |
+
│ ├── grader.py
|
| 93 |
+
│ └── __init__.py
|
| 94 |
+
└── next_quarter_move/ # 🚧 Stub
|
| 95 |
+
├── spec.py
|
| 96 |
+
├── grader.py
|
| 97 |
+
└── __init__.py
|
| 98 |
+
```
|
| 99 |
+
|
| 100 |
+
---
|
| 101 |
+
|
| 102 |
+
## 3. Architecture & Data Flow
|
| 103 |
+
|
| 104 |
+
```mermaid
|
| 105 |
+
sequenceDiagram
|
| 106 |
+
participant Script as inference.py / evaluate.py
|
| 107 |
+
participant Client as EarningsAnalystEnv (client.py)
|
| 108 |
+
participant Server as FastAPI (server/app.py)
|
| 109 |
+
participant Env as EarningsAnalystEnvironment
|
| 110 |
+
participant DS as HF Dataset (parquet singleton)
|
| 111 |
+
participant Registry as tasks/registry.py
|
| 112 |
+
|
| 113 |
+
Script->>Client: async with EarningsAnalystEnv(base_url)
|
| 114 |
+
Client->>Server: WS /ws — connect
|
| 115 |
+
|
| 116 |
+
Script->>Client: await env.reset()
|
| 117 |
+
Client->>Server: {type: "reset"}
|
| 118 |
+
Server->>Env: reset()
|
| 119 |
+
Env->>Registry: TASKS[task_id] → TaskSpec
|
| 120 |
+
Env->>DS: dataset[random_idx]
|
| 121 |
+
Env-->>Server: EarningsAnalystObservation (text_context, numerical_context, task_instruction)
|
| 122 |
+
Server-->>Client: JSON payload
|
| 123 |
+
Client-->>Script: StepResult.observation
|
| 124 |
+
|
| 125 |
+
Script->>Client: await env.step(EarningsAnalystAction(prediction="neutral"))
|
| 126 |
+
Client->>Server: {type: "step", prediction: "neutral"}
|
| 127 |
+
Server->>Env: step(action)
|
| 128 |
+
Env->>Registry: get_grader(task_id) → grade()
|
| 129 |
+
Env->>Env: reward = grade(predicted, ground_truth, label_values)
|
| 130 |
+
Env-->>Server: EarningsAnalystObservation (done=True, reward=0.5, metadata={...})
|
| 131 |
+
Server-->>Client: JSON payload
|
| 132 |
+
Client-->>Script: StepResult(reward=0.5, done=True)
|
| 133 |
+
```
|
| 134 |
+
|
| 135 |
+
**Key transport details:**
|
| 136 |
+
- The client uses a **persistent WebSocket** (`WS /ws`), not HTTP POST per call, for lower latency.
|
| 137 |
+
- The server also exposes `POST /reset`, `POST /step`, `GET /state`, and `GET /schema` for HTTP-only clients.
|
| 138 |
+
- `max_concurrent_envs=1` in `app.py` — increase for parallel evaluation runs.
|
| 139 |
+
|
| 140 |
+
---
|
| 141 |
+
|
| 142 |
+
## 4. Module Reference
|
| 143 |
+
|
| 144 |
+
### 4.1 Root Package (`earnings_analyst`)
|
| 145 |
+
|
| 146 |
+
**`__init__.py`** — Public surface of the installable package:
|
| 147 |
+
|
| 148 |
+
```python
|
| 149 |
+
from .client import EarningsAnalystEnv
|
| 150 |
+
from .models import EarningsAnalystAction, EarningsAnalystObservation
|
| 151 |
+
|
| 152 |
+
__all__ = ["EarningsAnalystAction", "EarningsAnalystObservation", "EarningsAnalystEnv"]
|
| 153 |
+
```
|
| 154 |
+
|
| 155 |
+
External consumers import these three names; everything else is internal.
|
| 156 |
+
|
| 157 |
+
---
|
| 158 |
+
|
| 159 |
+
### 4.2 `models.py`
|
| 160 |
+
|
| 161 |
+
Defines the two Pydantic models that flow between client and server.
|
| 162 |
+
|
| 163 |
+
#### `EarningsAnalystAction(Action)`
|
| 164 |
+
|
| 165 |
+
| Field | Type | Description |
|
| 166 |
+
|-------|------|-------------|
|
| 167 |
+
| `prediction` | `str` | Agent's answer — format depends on task (e.g. `"bullish"`, `"0.032"`). |
|
| 168 |
+
|
| 169 |
+
#### `EarningsAnalystObservation(Observation)`
|
| 170 |
+
|
| 171 |
+
| Field | Type | Description |
|
| 172 |
+
|-------|------|-------------|
|
| 173 |
+
| `text_context` | `dict[str, str]` | Keyed by column name. Only columns listed in `spec["text_cols"]` with non-empty values appear here. |
|
| 174 |
+
| `numerical_context` | `dict[str, float]` | Keyed by column name. Only finite (non-NaN) values from `spec["numerical_cols"]` appear. |
|
| 175 |
+
| `task_instruction` | `str` | Natural-language prompt copied from `spec["task_instruction"]`. Tells the agent exactly what to return. |
|
| 176 |
+
| `done` | `bool` | `False` after `reset()`, `True` after `step()`. |
|
| 177 |
+
| `reward` | `float \| None` | `0.0` after `reset()`, task-specific score after `step()`. |
|
| 178 |
+
| `metadata` | `dict` | After `step()`: `{task_id, predicted, ground_truth}`. Empty after `reset()`. |
|
| 179 |
+
|
| 180 |
+
---
|
| 181 |
+
|
| 182 |
+
### 4.3 `client.py`
|
| 183 |
+
|
| 184 |
+
**`EarningsAnalystEnv(EnvClient[Action, Observation, State])`**
|
| 185 |
+
|
| 186 |
+
A typed wrapper around `openenv-core`'s `EnvClient`. Used as an async context manager.
|
| 187 |
+
|
| 188 |
+
| Method | Purpose |
|
| 189 |
+
|--------|---------|
|
| 190 |
+
| `_step_payload(action)` | Serializes action to `{"prediction": str}` for the WebSocket message. |
|
| 191 |
+
| `_parse_result(payload)` | Deserializes JSON response into `StepResult[EarningsAnalystObservation]`. |
|
| 192 |
+
| `_parse_state(payload)` | Deserializes `{"episode_id", "step_count"}` into `State`. |
|
| 193 |
+
|
| 194 |
+
> [!NOTE]
|
| 195 |
+
> `EarningsAnalystEnv` also supports `from_docker_image(image_tag)` — a class method from the base class that spins up a Docker container and connects automatically.
|
| 196 |
+
|
| 197 |
+
**Typical usage:**
|
| 198 |
+
```python
|
| 199 |
+
async with EarningsAnalystEnv(base_url="http://localhost:8000") as env:
|
| 200 |
+
reset_result = await env.reset()
|
| 201 |
+
obs = reset_result.observation # EarningsAnalystObservation
|
| 202 |
+
step_result = await env.step(EarningsAnalystAction(prediction="bullish"))
|
| 203 |
+
print(step_result.reward)
|
| 204 |
+
```
|
| 205 |
+
|
| 206 |
+
---
|
| 207 |
+
|
| 208 |
+
### 4.4 `environment_config.py`
|
| 209 |
+
|
| 210 |
+
Thin config module. Provides two constants and re-exports task registry symbols.
|
| 211 |
+
|
| 212 |
+
| Symbol | Value |
|
| 213 |
+
|--------|-------|
|
| 214 |
+
| `DATASET_ID` | `"RudrakshNanavaty/earnings-call-data"` |
|
| 215 |
+
| `DATASET_FILE` | `"episodes_press_release_8k.parquet"` |
|
| 216 |
+
| `DEFAULT_TASK` | `"sentiment_label"` (re-exported from registry) |
|
| 217 |
+
| `TASKS` | `dict[str, TaskSpec]` (re-exported) |
|
| 218 |
+
|
| 219 |
+
---
|
| 220 |
+
|
| 221 |
+
### 4.5 `server/`
|
| 222 |
+
|
| 223 |
+
#### `server/dataset_loader.py`
|
| 224 |
+
|
| 225 |
+
Executes **once on first import** (module-level singleton):
|
| 226 |
+
|
| 227 |
+
```python
|
| 228 |
+
dataset = load_dataset(
|
| 229 |
+
DATASET_ID,
|
| 230 |
+
data_files={"train": DATASET_FILE},
|
| 231 |
+
split="train",
|
| 232 |
+
)
|
| 233 |
+
```
|
| 234 |
+
|
| 235 |
+
- Pins the specific parquet file to avoid silently picking up other files in the same Hub repo.
|
| 236 |
+
- All `reset()` calls reference this single in-memory object; no repeated downloads.
|
| 237 |
+
- On first run, `datasets` will download and cache to `~/.cache/huggingface/`.
|
| 238 |
+
|
| 239 |
+
#### `server/app.py`
|
| 240 |
+
|
| 241 |
+
Creates the FastAPI application via OpenEnv's factory:
|
| 242 |
+
|
| 243 |
+
```python
|
| 244 |
+
app = create_app(
|
| 245 |
+
EarningsAnalystEnvironment,
|
| 246 |
+
EarningsAnalystAction,
|
| 247 |
+
EarningsAnalystObservation,
|
| 248 |
+
env_name="earnings_analyst",
|
| 249 |
+
max_concurrent_envs=1,
|
| 250 |
+
)
|
| 251 |
+
```
|
| 252 |
+
|
| 253 |
+
**`main(host, port)`** — entrypoint used by `uv run server` (defined in `pyproject.toml` as `server = "earnings_analyst.server.app:main"`). Accepts `--port` via argparse.
|
| 254 |
+
|
| 255 |
+
Exposed HTTP/WS endpoints (from OpenEnv):
|
| 256 |
+
|
| 257 |
+
| Endpoint | Method | Purpose |
|
| 258 |
+
|----------|--------|---------|
|
| 259 |
+
| `/reset` | POST | Reset and return initial observation |
|
| 260 |
+
| `/step` | POST | Execute action, receive reward |
|
| 261 |
+
| `/state` | GET | Current `episode_id` + `step_count` |
|
| 262 |
+
| `/schema` | GET | JSON schemas for Action and Observation |
|
| 263 |
+
| `/ws` | WebSocket | Persistent session (preferred by client) |
|
| 264 |
+
| `/health` | GET | Docker health check target |
|
| 265 |
+
| `/web` | GET | OpenEnv web UI (base_path in HF Spaces) |
|
| 266 |
+
|
| 267 |
+
#### `server/earnings_analyst_environment.py`
|
| 268 |
+
|
| 269 |
+
**`EarningsAnalystEnvironment(Environment)`** — the core RL logic.
|
| 270 |
+
|
| 271 |
+
```
|
| 272 |
+
__init__(task_id=None)
|
| 273 |
+
↳ _resolve_task_id() # env var → DEFAULT_TASK fallback
|
| 274 |
+
↳ TASKS[task_id] # KeyError if unknown
|
| 275 |
+
↳ State(episode_id=uuid4)
|
| 276 |
+
|
| 277 |
+
reset() → EarningsAnalystObservation
|
| 278 |
+
↳ Check cfg["implemented"] or raise TaskNotImplementedError
|
| 279 |
+
↳ random row from dataset
|
| 280 |
+
↳ Filter text_cols (non-empty strings only)
|
| 281 |
+
↳ Filter numerical_cols (finite floats only, drops NaN)
|
| 282 |
+
↳ Return observation (done=False, reward=0.0)
|
| 283 |
+
|
| 284 |
+
step(action) → EarningsAnalystObservation
|
| 285 |
+
↳ get_grader(task_id)(prediction, ground_truth, label_values)
|
| 286 |
+
↳ Return terminal observation (done=True, reward=float, metadata)
|
| 287 |
+
|
| 288 |
+
state → State (property)
|
| 289 |
+
```
|
| 290 |
+
|
| 291 |
+
**Helper functions in the same file:**
|
| 292 |
+
|
| 293 |
+
| Function | Purpose |
|
| 294 |
+
|----------|---------|
|
| 295 |
+
| `_resolve_task_id(explicit)` | `explicit` → `EARNINGS_ANALYST_TASK_ID` env var → `DEFAULT_TASK` |
|
| 296 |
+
| `_non_empty_text(value)` | Returns `True` if value is a non-blank string |
|
| 297 |
+
| `_finite_float(value)` | Converts to float; returns `None` for NaN or unconvertible values |
|
| 298 |
+
|
| 299 |
+
> [!IMPORTANT]
|
| 300 |
+
> `SUPPORTS_CONCURRENT_SESSIONS = True` is set on the class, meaning the server can manage multiple independent environments if `max_concurrent_envs` in `app.py` is increased.
|
| 301 |
+
|
| 302 |
+
---
|
| 303 |
+
|
| 304 |
+
### 4.6 `tasks/`
|
| 305 |
+
|
| 306 |
+
The task system is the main extension point. Full breakdown in [§5](#5-task-system-deep-dive).
|
| 307 |
+
|
| 308 |
+
| Module | Exports |
|
| 309 |
+
|--------|---------|
|
| 310 |
+
| `tasks/types.py` | `TaskSpec` (TypedDict) |
|
| 311 |
+
| `tasks/exceptions.py` | `TaskNotImplementedError` |
|
| 312 |
+
| `tasks/grading.py` | `grade_ordinal()`, `grade_exact()` |
|
| 313 |
+
| `tasks/loader.py` | `load_task_subpackage()` |
|
| 314 |
+
| `tasks/registry.py` | `TASKS`, `GRADERS`, `TASK_IDS`, `DEFAULT_TASK`, `get_task_spec()`, `get_grader()` |
|
| 315 |
+
| `tasks/__init__.py` | Re-exports all of the above |
|
| 316 |
+
|
| 317 |
+
---
|
| 318 |
+
|
| 319 |
+
## 5. Task System Deep-Dive
|
| 320 |
+
|
| 321 |
+
### 5.1 `TaskSpec` schema
|
| 322 |
+
|
| 323 |
+
Defined in `tasks/types.py` as a `TypedDict`:
|
| 324 |
+
|
| 325 |
+
| Field | Type | Required | Description |
|
| 326 |
+
|-------|------|----------|-------------|
|
| 327 |
+
| `task_id` | `str` | ✅ | Unique identifier (e.g. `"sentiment_label"`). Must match the key in `TASKS`. |
|
| 328 |
+
| `implemented` | `bool` | ✅ | Gate flag. `False` → `reset()` raises `TaskNotImplementedError`. |
|
| 329 |
+
| `text_cols` | `list[str]` | ✅ | Dataset column names to include as `text_context`. |
|
| 330 |
+
| `numerical_cols` | `list[str]` | ✅ | Dataset column names to include as `numerical_context`. |
|
| 331 |
+
| `label_col` | `str` | ✅ | Column used as ground truth during `step()`. |
|
| 332 |
+
| `label_values` | `list[str]` | ✅ | Ordered list of valid labels (used by ordinal graders; empty for regression stubs). |
|
| 333 |
+
| `task_instruction` | `str` | ✅ | Full natural-language prompt shown to the agent as `observation.task_instruction`. |
|
| 334 |
+
| `kind` | `Literal["classification", "regression", "other"]` | ✅ | Metadata for reporting; not enforced by the environment. |
|
| 335 |
+
|
| 336 |
+
### 5.2 Task registry
|
| 337 |
+
|
| 338 |
+
**`tasks/registry.py`** is the single source of truth for all tasks.
|
| 339 |
+
|
| 340 |
+
```python
|
| 341 |
+
_TASK_ENTRIES: list[tuple[TaskSpec, GradingFn]] = [
|
| 342 |
+
(sentiment_label.SPEC, sentiment_label.grade),
|
| 343 |
+
(_pkg_1_day_move.SPEC, _pkg_1_day_move.grade),
|
| 344 |
+
(_pkg_30_day_move.SPEC, _pkg_30_day_move.grade),
|
| 345 |
+
(next_quarter_move.SPEC, next_quarter_move.grade),
|
| 346 |
+
]
|
| 347 |
+
|
| 348 |
+
TASKS: dict[str, TaskSpec] # keyed by task_id
|
| 349 |
+
GRADERS: dict[str, GradingFn] # keyed by task_id
|
| 350 |
+
DEFAULT_TASK = "sentiment_label"
|
| 351 |
+
```
|
| 352 |
+
|
| 353 |
+
`GradingFn` type alias: `Callable[[str, str, list[str]], float]` — `(predicted, ground_truth, label_values) → reward`.
|
| 354 |
+
|
| 355 |
+
**Why `load_task_subpackage()`?**
|
| 356 |
+
Python module names cannot start with a digit. Folders `1_day_move` and `30_day_move` are not importable via `import tasks.1_day_move`. `loader.py` uses `importlib.util.spec_from_file_location` to load them under synthetic qualified names (`earnings_analyst.tasks._pkg_1_day_move`, etc.) and injects them into `sys.modules`.
|
| 357 |
+
|
| 358 |
+
### 5.3 Grading helpers
|
| 359 |
+
|
| 360 |
+
**`tasks/grading.py`**:
|
| 361 |
+
|
| 362 |
+
#### `grade_ordinal(predicted, ground_truth, label_values) → float`
|
| 363 |
+
|
| 364 |
+
Ordinal (rank-aware) similarity reward for ordered label lists.
|
| 365 |
+
|
| 366 |
+
| Condition | Reward |
|
| 367 |
+
|-----------|--------|
|
| 368 |
+
| Exact match | `1.0` |
|
| 369 |
+
| Adjacent label (distance = 1) | `0.5` |
|
| 370 |
+
| Distance ≥ 2 or label not found | `0.0` |
|
| 371 |
+
|
| 372 |
+
Example for `label_values = ["very bearish", "bearish", "neutral", "bullish", "very bullish"]`:
|
| 373 |
+
- predicted=`"bullish"`, truth=`"bullish"` → `1.0`
|
| 374 |
+
- predicted=`"neutral"`, truth=`"bullish"` → `0.5`
|
| 375 |
+
- predicted=`"very bearish"`, truth=`"very bullish"` → `0.0`
|
| 376 |
+
|
| 377 |
+
#### `grade_exact(predicted, ground_truth, label_values) → float`
|
| 378 |
+
|
| 379 |
+
Binary reward: `1.0` if case-insensitive stripped match, else `0.0`. `label_values` is ignored.
|
| 380 |
+
|
| 381 |
+
### 5.4 Task inventory
|
| 382 |
+
|
| 383 |
+
| `task_id` | `kind` | Status | Grader | Text columns | Numerical columns |
|
| 384 |
+
|-----------|--------|--------|--------|--------------|-------------------|
|
| 385 |
+
| `sentiment_label` | classification | ✅ Implemented | `grade_ordinal` | `earnings_transcript`, `press_release_8k_body`, `press_release_ex991`, `press_release_ex992` | `price_momentum_30d`, `price_momentum_90d`, `pct_from_52w_high_pt`, `avg_volume_20d`, `d_minus_1_close` |
|
| 386 |
+
| `1_day_move` | regression | 🚧 Stub | raises `NotImplementedError` | (none) | (none) |
|
| 387 |
+
| `30_day_move` | regression | 🚧 Stub | raises `NotImplementedError` | (none) | (none) |
|
| 388 |
+
| `next_quarter_move` | regression | 🚧 Stub | raises `NotImplementedError` | (none) | (none) |
|
| 389 |
+
|
| 390 |
+
**`sentiment_label` label order** (ordinal, worst → best):
|
| 391 |
+
```
|
| 392 |
+
very bearish → bearish → neutral → bullish → very bullish
|
| 393 |
+
```
|
| 394 |
+
|
| 395 |
+
---
|
| 396 |
+
|
| 397 |
+
## 6. Environment Lifecycle
|
| 398 |
+
|
| 399 |
+
```
|
| 400 |
+
Server startup
|
| 401 |
+
│
|
| 402 |
+
├── dataset_loader.py imported → HF parquet fetched/cached → dataset singleton
|
| 403 |
+
└── EarningsAnalystEnvironment.__init__()
|
| 404 |
+
├── _resolve_task_id() (env var / default)
|
| 405 |
+
├── TASKS[task_id] (KeyError if unknown)
|
| 406 |
+
└── State(episode_id=uuid4, step_count=0)
|
| 407 |
+
│
|
| 408 |
+
▼ client calls reset()
|
| 409 |
+
reset()
|
| 410 |
+
├── Check cfg["implemented"] → TaskNotImplementedError if False
|
| 411 |
+
├── State reset (new uuid4, step_count=0)
|
| 412 |
+
├── dataset[random.randrange(len(dataset))]
|
| 413 |
+
├── Build text_context (non-empty strings from text_cols)
|
| 414 |
+
├── Build numerical_context (finite floats from numerical_cols)
|
| 415 |
+
└── Return EarningsAnalystObservation(done=False, reward=0.0)
|
| 416 |
+
│
|
| 417 |
+
▼ client calls step(action)
|
| 418 |
+
step(action)
|
| 419 |
+
├── state.step_count += 1
|
| 420 |
+
├── ground_truth = row[label_col]
|
| 421 |
+
├── grade_fn = get_grader(task_id)
|
| 422 |
+
├── reward = float(grade_fn(prediction, ground_truth, label_values))
|
| 423 |
+
└── Return EarningsAnalystObservation(done=True, reward=reward,
|
| 424 |
+
metadata={task_id, predicted, ground_truth})
|
| 425 |
+
```
|
| 426 |
+
|
| 427 |
+
> [!NOTE]
|
| 428 |
+
> Each episode is **always terminal after exactly one step**. There is no multi-step trajectory — the environment is episodic/bandit-style.
|
| 429 |
+
|
| 430 |
+
---
|
| 431 |
+
|
| 432 |
+
## 7. Inference & Evaluation Scripts
|
| 433 |
+
|
| 434 |
+
### `inference.py`
|
| 435 |
+
|
| 436 |
+
Single-episode example implementation using OpenAI Chat Completions. Currently **hardcoded for sentiment classification** — requires adaptation for other tasks.
|
| 437 |
+
|
| 438 |
+
**Key functions:**
|
| 439 |
+
|
| 440 |
+
| Function | Signature | Purpose |
|
| 441 |
+
|----------|-----------|---------|
|
| 442 |
+
| `build_user_content(obs)` | `EarningsAnalystObservation → str` | Assembles the user message from `task_instruction` + formatted `text_context` + JSON `numerical_context`. |
|
| 443 |
+
| `predict_with_openai(obs, *, client, model, valid_labels)` | `→ tuple[str, str]` | Calls Chat Completions with `response_format=json_object`, parses `{"sentiment": ...}`, normalizes with `_normalize_sentiment()`. |
|
| 444 |
+
| `_normalize_sentiment(text, valid)` | `→ str` | Case-insensitive fuzzy match to canonical labels; fallback `"neutral"`. |
|
| 445 |
+
| `run_episode(*, base_url, openai_api_key, ...)` | `async → EpisodeResult` | Full `reset → predict → step` flow. Returns `EpisodeResult(reward, predicted, ground_truth, done, model_response_text)`. |
|
| 446 |
+
| `main()` | CLI entrypoint | Parses `--base-url`, `--model`, `--quiet`. |
|
| 447 |
+
|
| 448 |
+
**CLI:**
|
| 449 |
+
```bash
|
| 450 |
+
uv run python inference.py
|
| 451 |
+
uv run python inference.py --base-url http://localhost:8000 --model gpt-4o-mini --quiet
|
| 452 |
+
```
|
| 453 |
+
|
| 454 |
+
---
|
| 455 |
+
|
| 456 |
+
### `evaluate.py`
|
| 457 |
+
|
| 458 |
+
Runs `N` independent episodes (each a full `run_episode()` call) and aggregates metrics.
|
| 459 |
+
|
| 460 |
+
**Key functions:**
|
| 461 |
+
|
| 462 |
+
| Function | Purpose |
|
| 463 |
+
|----------|---------|
|
| 464 |
+
| `exact_match(predicted, ground_truth)` | Case-insensitive strip comparison → bool |
|
| 465 |
+
| `confusion_key(predicted, ground_truth)` | Returns normalized `(predicted, truth)` tuple for confusion matrix |
|
| 466 |
+
| `run_evaluation(*, samples, base_url, model, task_id, quiet)` | Main async loop; prints summary table |
|
| 467 |
+
|
| 468 |
+
**Output:**
|
| 469 |
+
```
|
| 470 |
+
=== Evaluation summary ===
|
| 471 |
+
samples: 100
|
| 472 |
+
mean_reward: 0.6250
|
| 473 |
+
exact_accuracy: 0.5000 (50/100)
|
| 474 |
+
|
| 475 |
+
Per ground-truth label (exact match rate):
|
| 476 |
+
'very bearish': 0.4000 (2/5)
|
| 477 |
+
'bearish': 0.3333 (4/12)
|
| 478 |
+
...
|
| 479 |
+
|
| 480 |
+
Confusion (predicted -> counts by ground_truth):
|
| 481 |
+
truth='bullish': 'bullish':30, 'neutral':5
|
| 482 |
+
...
|
| 483 |
+
```
|
| 484 |
+
|
| 485 |
+
**CLI:**
|
| 486 |
+
```bash
|
| 487 |
+
uv run python evaluate.py
|
| 488 |
+
uv run python evaluate.py --samples 50 --task sentiment_label --quiet
|
| 489 |
+
```
|
| 490 |
+
|
| 491 |
+
> [!WARNING]
|
| 492 |
+
> The `--task` flag only controls which `TaskSpec` is used for **reporting** (label list). It does **not** change the server's active task. The server's `EARNINGS_ANALYST_TASK_ID` must be set before startup.
|
| 493 |
+
|
| 494 |
+
---
|
| 495 |
+
|
| 496 |
+
## 8. Configuration Reference
|
| 497 |
+
|
| 498 |
+
### Environment variables
|
| 499 |
+
|
| 500 |
+
| Variable | Required | Default | Description |
|
| 501 |
+
|----------|----------|---------|-------------|
|
| 502 |
+
| `OPENAI_API_KEY` | For inference/evaluate | — | Chat Completions key |
|
| 503 |
+
| `OPENAI_BASE_URL` | No | OpenAI default | Custom API base (proxies, Azure, Google OpenAI-compat) |
|
| 504 |
+
| `OPENAI_MODEL` | No | `gpt-4o` | Model ID for inference scripts |
|
| 505 |
+
| `ENV_SERVER_URL` | No | `http://localhost:8000` | Base URL for `EarningsAnalystEnv` in client scripts |
|
| 506 |
+
| `EARNINGS_ANALYST_TASK_ID` | No | `sentiment_label` | Task loaded at **server startup** |
|
| 507 |
+
| `HF_TOKEN` | For private datasets | — | Hugging Face auth token |
|
| 508 |
+
|
| 509 |
+
Load order: `python-dotenv` reads `.env` first (copy from `.env.example`).
|
| 510 |
+
|
| 511 |
+
### Dataset configuration (`environment_config.py`)
|
| 512 |
+
|
| 513 |
+
| Constant | Value |
|
| 514 |
+
|----------|-------|
|
| 515 |
+
| `DATASET_ID` | `RudrakshNanavaty/earnings-call-data` |
|
| 516 |
+
| `DATASET_FILE` | `episodes_press_release_8k.parquet` |
|
| 517 |
+
|
| 518 |
+
To switch datasets: edit these two constants and restart the server.
|
| 519 |
+
|
| 520 |
+
---
|
| 521 |
+
|
| 522 |
+
## 9. Deployment
|
| 523 |
+
|
| 524 |
+
### Local development
|
| 525 |
+
|
| 526 |
+
```bash
|
| 527 |
+
uv sync # install deps
|
| 528 |
+
uv run server # starts on 0.0.0.0:8000
|
| 529 |
+
uv run server --port 8001 # custom port
|
| 530 |
+
|
| 531 |
+
# Or with uvicorn directly (auto-reload):
|
| 532 |
+
uv run uvicorn server.app:app --host 0.0.0.0 --port 8000 --reload
|
| 533 |
+
```
|
| 534 |
+
|
| 535 |
+
Set task before starting:
|
| 536 |
+
```bash
|
| 537 |
+
export EARNINGS_ANALYST_TASK_ID=sentiment_label
|
| 538 |
+
uv run server
|
| 539 |
+
```
|
| 540 |
+
|
| 541 |
+
### Docker (`server/Dockerfile`)
|
| 542 |
+
|
| 543 |
+
```bash
|
| 544 |
+
# Build
|
| 545 |
+
docker build -t earnings_analyst-env:latest -f server/Dockerfile .
|
| 546 |
+
|
| 547 |
+
# Run
|
| 548 |
+
docker run -p 8000:8000 \
|
| 549 |
+
-e EARNINGS_ANALYST_TASK_ID=sentiment_label \
|
| 550 |
+
-e HF_TOKEN=hf_... \
|
| 551 |
+
earnings_analyst-env:latest
|
| 552 |
+
```
|
| 553 |
+
|
| 554 |
+
The Dockerfile:
|
| 555 |
+
- Uses `openenv-base` as the base image
|
| 556 |
+
- Runs `uv sync` to install deps
|
| 557 |
+
- Entry: `uvicorn server.app:app` on port 8000
|
| 558 |
+
- Health check: `GET /health`
|
| 559 |
+
|
| 560 |
+
### Hugging Face Spaces
|
| 561 |
+
|
| 562 |
+
The repo frontmatter and `openenv.yaml` configure HF Spaces deployment:
|
| 563 |
+
- `sdk: docker` → uses `server/Dockerfile`
|
| 564 |
+
- `app_port: 8000`, `base_path: /web`
|
| 565 |
+
- App entry: `server.app:app`
|
| 566 |
+
|
| 567 |
+
Pass `EARNINGS_ANALYST_TASK_ID` and `HF_TOKEN` as Space secrets.
|
| 568 |
+
|
| 569 |
+
---
|
| 570 |
+
|
| 571 |
+
## 10. Adding a New Task
|
| 572 |
+
|
| 573 |
+
### Step-by-step
|
| 574 |
+
|
| 575 |
+
**1. Create the task directory**
|
| 576 |
+
|
| 577 |
+
```
|
| 578 |
+
tasks/
|
| 579 |
+
└── my_task/
|
| 580 |
+
├── __init__.py
|
| 581 |
+
├── spec.py
|
| 582 |
+
└── grader.py
|
| 583 |
+
```
|
| 584 |
+
|
| 585 |
+
> [!TIP]
|
| 586 |
+
> If the folder name starts with a digit (e.g. `5_day_move`), follow the `1_day_move` pattern and use `load_task_subpackage` in `registry.py`.
|
| 587 |
+
|
| 588 |
+
**2. Define `spec.py`**
|
| 589 |
+
|
| 590 |
+
```python
|
| 591 |
+
from ..types import TaskSpec
|
| 592 |
+
|
| 593 |
+
CANONICAL_TASK_ID = "my_task"
|
| 594 |
+
|
| 595 |
+
SPEC: TaskSpec = {
|
| 596 |
+
"task_id": CANONICAL_TASK_ID,
|
| 597 |
+
"implemented": True, # Set False until ready
|
| 598 |
+
"text_cols": ["earnings_transcript"], # Columns from the parquet
|
| 599 |
+
"numerical_cols": ["price_momentum_30d"],
|
| 600 |
+
"label_col": "my_target_column", # Ground truth column
|
| 601 |
+
"label_values": ["low", "medium", "high"], # Ordered if ordinal
|
| 602 |
+
"task_instruction": (
|
| 603 |
+
"Predict the price category.\n\n"
|
| 604 |
+
'Return JSON: {"category": "<low|medium|high>"}'
|
| 605 |
+
),
|
| 606 |
+
"kind": "classification", # or "regression" / "other"
|
| 607 |
+
}
|
| 608 |
+
```
|
| 609 |
+
|
| 610 |
+
**3. Implement `grader.py`**
|
| 611 |
+
|
| 612 |
+
```python
|
| 613 |
+
from ..grading import grade_ordinal # or grade_exact, or custom logic
|
| 614 |
+
|
| 615 |
+
def grade(predicted: str, ground_truth: str, label_values: list[str]) -> float:
|
| 616 |
+
return grade_ordinal(predicted, ground_truth, label_values)
|
| 617 |
+
```
|
| 618 |
+
|
| 619 |
+
**4. Export from `__init__.py`**
|
| 620 |
+
|
| 621 |
+
```python
|
| 622 |
+
from .spec import SPEC
|
| 623 |
+
from .grader import grade
|
| 624 |
+
|
| 625 |
+
__all__ = ["SPEC", "grade"]
|
| 626 |
+
```
|
| 627 |
+
|
| 628 |
+
**5. Register in `tasks/registry.py`**
|
| 629 |
+
|
| 630 |
+
For a regular folder name:
|
| 631 |
+
```python
|
| 632 |
+
from . import my_task # add this import
|
| 633 |
+
|
| 634 |
+
_TASK_ENTRIES: list[...] = [
|
| 635 |
+
...
|
| 636 |
+
(my_task.SPEC, my_task.grade), # add this line
|
| 637 |
+
]
|
| 638 |
+
```
|
| 639 |
+
|
| 640 |
+
For a digit-prefixed folder name:
|
| 641 |
+
```python
|
| 642 |
+
_pkg_5_day_move = load_task_subpackage(
|
| 643 |
+
"5_day_move",
|
| 644 |
+
"earnings_analyst.tasks._pkg_5_day_move",
|
| 645 |
+
)
|
| 646 |
+
# then add to _TASK_ENTRIES:
|
| 647 |
+
(_pkg_5_day_move.SPEC, _pkg_5_day_move.grade),
|
| 648 |
+
```
|
| 649 |
+
|
| 650 |
+
**6. (If needed) Update `pyproject.toml`**
|
| 651 |
+
|
| 652 |
+
Regular folders are already covered by `package-dir`. For digit-prefixed folders, they're loaded as data files — confirm `package-data` includes `"my_task/*.py"`:
|
| 653 |
+
```toml
|
| 654 |
+
[tool.setuptools.package-data]
|
| 655 |
+
"earnings_analyst.tasks" = ["1_day_move/*.py", "30_day_move/*.py", "5_day_move/*.py"]
|
| 656 |
+
```
|
| 657 |
+
|
| 658 |
+
**7. Restart the server**
|
| 659 |
+
|
| 660 |
+
```bash
|
| 661 |
+
export EARNINGS_ANALYST_TASK_ID=my_task
|
| 662 |
+
uv run server
|
| 663 |
+
```
|
| 664 |
+
|
| 665 |
+
---
|
| 666 |
+
|
| 667 |
+
## 11. Key Design Decisions & Gotchas
|
| 668 |
+
|
| 669 |
+
### Dual-import fallback pattern
|
| 670 |
+
|
| 671 |
+
Every module uses a try/except import pattern:
|
| 672 |
+
```python
|
| 673 |
+
try:
|
| 674 |
+
from earnings_analyst.tasks.registry import ... # installed package
|
| 675 |
+
except ImportError:
|
| 676 |
+
from tasks.registry import ... # repo root on PYTHONPATH
|
| 677 |
+
```
|
| 678 |
+
This makes scripts runnable both as `uv run python inference.py` (PYTHONPATH=root) and from the installed package.
|
| 679 |
+
|
| 680 |
+
### Dataset is a global singleton
|
| 681 |
+
|
| 682 |
+
`dataset_loader.py` loads the HF dataset **once** at import time. Benefits: fast resets, no redundant downloads. Caveat: the server process must have sufficient memory to hold the dataset. There is no streaming or lazy loading.
|
| 683 |
+
|
| 684 |
+
### Single-step episodes (bandit)
|
| 685 |
+
|
| 686 |
+
Every episode is exactly `reset() → step()`. `done=True` is always returned from `step()`. This means standard RL algorithms that expect trajectories need adaptation (treat as contextual bandit).
|
| 687 |
+
|
| 688 |
+
### `max_concurrent_envs=1`
|
| 689 |
+
|
| 690 |
+
The server currently supports only one concurrent WebSocket session. For parallel evaluation, either:
|
| 691 |
+
- Run multiple server processes on different ports, or
|
| 692 |
+
- Increase `max_concurrent_envs` in `server/app.py`.
|
| 693 |
+
|
| 694 |
+
### `inference.py` is sentiment-specific
|
| 695 |
+
|
| 696 |
+
The `predict_with_openai()` function is hardcoded to return `{"sentiment": ...}`. For other tasks, you must write a new prediction function (different system prompt, different JSON key, different normalization).
|
| 697 |
+
|
| 698 |
+
### Task server vs. evaluation script alignment
|
| 699 |
+
|
| 700 |
+
`evaluate.py --task <id>` only selects which `TaskSpec` to use for **printing per-label stats**. The actual grading happens on the **server** using `EARNINGS_ANALYST_TASK_ID`. If these disagree, label statistics will be meaningless (wrong label list).
|
| 701 |
+
|
| 702 |
+
### `implemented` gate is enforced only in `reset()`
|
| 703 |
+
|
| 704 |
+
Setting `implemented: False` prevents `reset()` but does **not** prevent the task from appearing in `TASKS` or `GRADERS`. The grader's `NotImplementedError` is only hit if someone calls `step()` on a row from a stub task — which can't happen via normal flow since `reset()` blocks first.
|
earnings_analyst/models.py
CHANGED
|
@@ -2,6 +2,7 @@
|
|
| 2 |
Data models for the Earnings Analyst Environment.
|
| 3 |
"""
|
| 4 |
|
|
|
|
| 5 |
from openenv.core.env_server.types import Action, Observation
|
| 6 |
from pydantic import Field
|
| 7 |
|
|
@@ -33,3 +34,7 @@ class EarningsAnalystObservation(Observation):
|
|
| 33 |
default="",
|
| 34 |
description="Natural language instruction and JSON schema for the agent",
|
| 35 |
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 2 |
Data models for the Earnings Analyst Environment.
|
| 3 |
"""
|
| 4 |
|
| 5 |
+
from typing import Any
|
| 6 |
from openenv.core.env_server.types import Action, Observation
|
| 7 |
from pydantic import Field
|
| 8 |
|
|
|
|
| 34 |
default="",
|
| 35 |
description="Natural language instruction and JSON schema for the agent",
|
| 36 |
)
|
| 37 |
+
metadata: dict[str, Any] = Field(
|
| 38 |
+
default_factory=dict,
|
| 39 |
+
description="Additional context for debugging, rewards, or logging (e.g. ground truth)",
|
| 40 |
+
)
|
evaluate.py
CHANGED
|
@@ -99,6 +99,17 @@ async def run_evaluation(
|
|
| 99 |
if is_exact:
|
| 100 |
per_ground_truth_label[normalized_ground_truth]["correct"] += 1
|
| 101 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 102 |
csv_rows.append(
|
| 103 |
{
|
| 104 |
"sample_index": episode_index + 1,
|
|
|
|
| 99 |
if is_exact:
|
| 100 |
per_ground_truth_label[normalized_ground_truth]["correct"] += 1
|
| 101 |
|
| 102 |
+
# --- VERBOSE PRINTING BLOCK (Safe to remove) ---
|
| 103 |
+
# NOTE: This block is exclusively for result visibility in the console.
|
| 104 |
+
if not quiet:
|
| 105 |
+
print(f"\n--- Episode {episode_index + 1}/{samples} Summary ---")
|
| 106 |
+
print(f"Reward: {episode_reward:.4f}")
|
| 107 |
+
print(f"Predicted: {predicted_label}")
|
| 108 |
+
print(f"Ground Truth: {ground_truth_label}")
|
| 109 |
+
print(f"Model Response: {episode_result.model_response_text}")
|
| 110 |
+
print("-" * 40)
|
| 111 |
+
# ------------------------------------------------
|
| 112 |
+
|
| 113 |
csv_rows.append(
|
| 114 |
{
|
| 115 |
"sample_index": episode_index + 1,
|
models.py
CHANGED
|
@@ -2,6 +2,7 @@
|
|
| 2 |
Data models for the Earnings Analyst Environment.
|
| 3 |
"""
|
| 4 |
|
|
|
|
| 5 |
from openenv.core.env_server.types import Action, Observation
|
| 6 |
from pydantic import Field
|
| 7 |
|
|
@@ -33,3 +34,11 @@ class EarningsAnalystObservation(Observation):
|
|
| 33 |
default="",
|
| 34 |
description="Natural language instruction and JSON schema for the agent",
|
| 35 |
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 2 |
Data models for the Earnings Analyst Environment.
|
| 3 |
"""
|
| 4 |
|
| 5 |
+
from typing import Any
|
| 6 |
from openenv.core.env_server.types import Action, Observation
|
| 7 |
from pydantic import Field
|
| 8 |
|
|
|
|
| 34 |
default="",
|
| 35 |
description="Natural language instruction and JSON schema for the agent",
|
| 36 |
)
|
| 37 |
+
ground_truth: str = Field(
|
| 38 |
+
default="",
|
| 39 |
+
description="Actual value for the task (populated on terminal step)",
|
| 40 |
+
)
|
| 41 |
+
metadata: dict[str, Any] = Field(
|
| 42 |
+
default_factory=dict,
|
| 43 |
+
description="Additional context for debugging, rewards, or logging",
|
| 44 |
+
)
|
scratch/inspect_dataset.py
ADDED
|
@@ -0,0 +1,26 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from datasets import load_dataset
|
| 2 |
+
import os
|
| 3 |
+
|
| 4 |
+
DATASET_ID = "RudrakshNanavaty/earnings-call-data"
|
| 5 |
+
DATASET_FILE = "episodes_press_release_8k.parquet"
|
| 6 |
+
|
| 7 |
+
print(f"Loading dataset {DATASET_ID} file {DATASET_FILE}...")
|
| 8 |
+
dataset = load_dataset(
|
| 9 |
+
DATASET_ID,
|
| 10 |
+
data_files={"train": DATASET_FILE},
|
| 11 |
+
split="train",
|
| 12 |
+
)
|
| 13 |
+
|
| 14 |
+
print("\nColumns:")
|
| 15 |
+
print(dataset.column_names)
|
| 16 |
+
|
| 17 |
+
print("\nFirst row summary:")
|
| 18 |
+
row = dataset[0]
|
| 19 |
+
for k, v in row.items():
|
| 20 |
+
if v is not None:
|
| 21 |
+
val_str = str(v)
|
| 22 |
+
if len(val_str) > 100:
|
| 23 |
+
val_str = val_str[:100] + "..."
|
| 24 |
+
print(f"{k}: {val_str}")
|
| 25 |
+
else:
|
| 26 |
+
print(f"{k}: None")
|
tasks/1_day_move/grader.py
CHANGED
|
@@ -1,10 +1,12 @@
|
|
| 1 |
-
"""Grading for ``1_day_move``
|
| 2 |
|
| 3 |
from __future__ import annotations
|
|
|
|
| 4 |
|
| 5 |
|
| 6 |
def grade(predicted: str, ground_truth: str, label_values: list[str]) -> float:
|
| 7 |
-
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
|
|
|
|
|
|
| 1 |
+
"""Grading logic for ``1_day_move``."""
|
| 2 |
|
| 3 |
from __future__ import annotations
|
| 4 |
+
from ..grading import grade_smart_move
|
| 5 |
|
| 6 |
|
| 7 |
def grade(predicted: str, ground_truth: str, label_values: list[str]) -> float:
|
| 8 |
+
"""
|
| 9 |
+
Score the agent's prediction using the smart reward (Directional + Numerical).
|
| 10 |
+
This logic is handled centrally in grading.py to ensure consistency across movement tasks.
|
| 11 |
+
"""
|
| 12 |
+
return grade_smart_move(predicted, ground_truth, label_values)
|
tasks/1_day_move/spec.py
CHANGED
|
@@ -8,11 +8,41 @@ CANONICAL_TASK_ID = "1_day_move"
|
|
| 8 |
|
| 9 |
SPEC: TaskSpec = {
|
| 10 |
"task_id": CANONICAL_TASK_ID,
|
| 11 |
-
"implemented":
|
| 12 |
-
"text_cols": [
|
| 13 |
-
|
| 14 |
-
|
| 15 |
-
|
| 16 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 17 |
"kind": "regression",
|
| 18 |
}
|
|
|
|
|
|
| 8 |
|
| 9 |
SPEC: TaskSpec = {
|
| 10 |
"task_id": CANONICAL_TASK_ID,
|
| 11 |
+
"implemented": True,
|
| 12 |
+
"text_cols": [
|
| 13 |
+
"earnings_transcript",
|
| 14 |
+
"press_release_8k_body",
|
| 15 |
+
"press_release_ex991",
|
| 16 |
+
"press_release_ex992",
|
| 17 |
+
"press_release_sources",
|
| 18 |
+
],
|
| 19 |
+
"numerical_cols": [
|
| 20 |
+
"price_momentum_30d",
|
| 21 |
+
"price_momentum_90d",
|
| 22 |
+
"pct_from_52w_high_pt",
|
| 23 |
+
"avg_volume_20d",
|
| 24 |
+
"d_minus_1_close",
|
| 25 |
+
],
|
| 26 |
+
"label_col": "move_1d",
|
| 27 |
+
"label_values": [
|
| 28 |
+
"very bearish",
|
| 29 |
+
"bearish",
|
| 30 |
+
"neutral",
|
| 31 |
+
"bullish",
|
| 32 |
+
"very bullish",
|
| 33 |
+
],
|
| 34 |
+
"task_instruction": (
|
| 35 |
+
"Analyse the provided earnings call materials and market data to predict the stock price movement after 1 day.\n\n"
|
| 36 |
+
"Return a JSON object matching this exact schema:\n"
|
| 37 |
+
'{"percentage_move": <float>, "label": "<one of: very bearish | bearish | neutral | bullish | very bullish>"}\n\n'
|
| 38 |
+
"Brackets:\n"
|
| 39 |
+
"- > 7% negative: very bearish\n"
|
| 40 |
+
"- 1-7% negative: bearish\n"
|
| 41 |
+
"- -1% to +1%: neutral\n"
|
| 42 |
+
"- 1-7% positive: bullish\n"
|
| 43 |
+
"- > 7% positive: very bullish\n\n"
|
| 44 |
+
"Do not include any other keys or explanation."
|
| 45 |
+
),
|
| 46 |
"kind": "regression",
|
| 47 |
}
|
| 48 |
+
|
tasks/30_day_move/grader.py
CHANGED
|
@@ -1,10 +1,12 @@
|
|
| 1 |
-
"""Grading for ``30_day_move``
|
| 2 |
|
| 3 |
from __future__ import annotations
|
|
|
|
| 4 |
|
| 5 |
|
| 6 |
def grade(predicted: str, ground_truth: str, label_values: list[str]) -> float:
|
| 7 |
-
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
|
|
|
|
|
|
| 1 |
+
"""Grading logic for ``30_day_move``."""
|
| 2 |
|
| 3 |
from __future__ import annotations
|
| 4 |
+
from ..grading import grade_smart_move
|
| 5 |
|
| 6 |
|
| 7 |
def grade(predicted: str, ground_truth: str, label_values: list[str]) -> float:
|
| 8 |
+
"""
|
| 9 |
+
Score the agent's prediction using the smart reward (Directional + Numerical).
|
| 10 |
+
This logic is handled centrally in grading.py to ensure consistency across movement tasks.
|
| 11 |
+
"""
|
| 12 |
+
return grade_smart_move(predicted, ground_truth, label_values)
|
tasks/30_day_move/spec.py
CHANGED
|
@@ -8,11 +8,41 @@ CANONICAL_TASK_ID = "30_day_move"
|
|
| 8 |
|
| 9 |
SPEC: TaskSpec = {
|
| 10 |
"task_id": CANONICAL_TASK_ID,
|
| 11 |
-
"implemented":
|
| 12 |
-
"text_cols": [
|
| 13 |
-
|
| 14 |
-
|
| 15 |
-
|
| 16 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 17 |
"kind": "regression",
|
| 18 |
}
|
|
|
|
|
|
| 8 |
|
| 9 |
SPEC: TaskSpec = {
|
| 10 |
"task_id": CANONICAL_TASK_ID,
|
| 11 |
+
"implemented": True,
|
| 12 |
+
"text_cols": [
|
| 13 |
+
"earnings_transcript",
|
| 14 |
+
"press_release_8k_body",
|
| 15 |
+
"press_release_ex991",
|
| 16 |
+
"press_release_ex992",
|
| 17 |
+
"press_release_sources",
|
| 18 |
+
],
|
| 19 |
+
"numerical_cols": [
|
| 20 |
+
"price_momentum_30d",
|
| 21 |
+
"price_momentum_90d",
|
| 22 |
+
"pct_from_52w_high_pt",
|
| 23 |
+
"avg_volume_20d",
|
| 24 |
+
"d_minus_1_close",
|
| 25 |
+
],
|
| 26 |
+
"label_col": "move_30d",
|
| 27 |
+
"label_values": [
|
| 28 |
+
"very bearish",
|
| 29 |
+
"bearish",
|
| 30 |
+
"neutral",
|
| 31 |
+
"bullish",
|
| 32 |
+
"very bullish",
|
| 33 |
+
],
|
| 34 |
+
"task_instruction": (
|
| 35 |
+
"Analyse the provided earnings call materials and market data to predict the stock price movement after 30 days.\n\n"
|
| 36 |
+
"Return a JSON object matching this exact schema:\n"
|
| 37 |
+
'{"percentage_move": <float>, "label": "<one of: very bearish | bearish | neutral | bullish | very bullish>"}\n\n'
|
| 38 |
+
"Brackets:\n"
|
| 39 |
+
"- > 7% negative: very bearish\n"
|
| 40 |
+
"- 1-7% negative: bearish\n"
|
| 41 |
+
"- -1% to +1%: neutral\n"
|
| 42 |
+
"- 1-7% positive: bullish\n"
|
| 43 |
+
"- > 7% positive: very bullish\n\n"
|
| 44 |
+
"Do not include any other keys or explanation."
|
| 45 |
+
),
|
| 46 |
"kind": "regression",
|
| 47 |
}
|
| 48 |
+
|
tasks/get_figures/__init__.py
ADDED
|
@@ -0,0 +1,4 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from .spec import SPEC
|
| 2 |
+
from .grader import grade
|
| 3 |
+
|
| 4 |
+
__all__ = ["SPEC", "grade"]
|
tasks/get_figures/grader.py
ADDED
|
@@ -0,0 +1,14 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Grading logic for ``get_figures``."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
import json
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
def grade(predicted: str, ground_truth: str, label_values: list[str]) -> float:
|
| 8 |
+
"""
|
| 9 |
+
Score the agent's extraction performance.
|
| 10 |
+
|
| 11 |
+
Currently a stub until XBRL ground truth is provided.
|
| 12 |
+
Always returns 0.0 with implemented: False in spec.
|
| 13 |
+
"""
|
| 14 |
+
return 0.0
|
tasks/get_figures/spec.py
ADDED
|
@@ -0,0 +1,29 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Task specification for ``get_figures`` (Financial figure extraction)."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
from ..types import TaskSpec
|
| 5 |
+
|
| 6 |
+
CANONICAL_TASK_ID = "get_figures"
|
| 7 |
+
|
| 8 |
+
SPEC: TaskSpec = {
|
| 9 |
+
"task_id": CANONICAL_TASK_ID,
|
| 10 |
+
"implemented": False, # Safety gate: set to True once ground truth column is confirmed.
|
| 11 |
+
"text_cols": [
|
| 12 |
+
"earnings_transcript",
|
| 13 |
+
"press_release_8k_body",
|
| 14 |
+
"press_release_ex991",
|
| 15 |
+
"press_release_ex992",
|
| 16 |
+
"press_release_sources",
|
| 17 |
+
],
|
| 18 |
+
"numerical_cols": [],
|
| 19 |
+
"label_col": "symbol", # Placeholder
|
| 20 |
+
"label_values": [],
|
| 21 |
+
"task_instruction": (
|
| 22 |
+
"Extract key financial figures from the provided earnings call materials.\n\n"
|
| 23 |
+
"Return a JSON object matching this exact schema:\n"
|
| 24 |
+
'{"revenue": <float>, "net_income": <float>, "eps": <float>}\n\n'
|
| 25 |
+
"Use the currency specified in the documents. If a figure is not found, use null.\n"
|
| 26 |
+
"Do not include any other keys or explanation."
|
| 27 |
+
),
|
| 28 |
+
"kind": "other",
|
| 29 |
+
}
|
tasks/grading.py
CHANGED
|
@@ -1,6 +1,8 @@
|
|
| 1 |
"""Shared grading helpers for task modules."""
|
| 2 |
|
| 3 |
from __future__ import annotations
|
|
|
|
|
|
|
| 4 |
|
| 5 |
|
| 6 |
def _normalize_text(text: str) -> str:
|
|
@@ -33,6 +35,56 @@ def grade_ordinal(
|
|
| 33 |
return 0.0
|
| 34 |
|
| 35 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 36 |
def grade_exact(
|
| 37 |
predicted: str,
|
| 38 |
ground_truth: str,
|
|
|
|
| 1 |
"""Shared grading helpers for task modules."""
|
| 2 |
|
| 3 |
from __future__ import annotations
|
| 4 |
+
import json
|
| 5 |
+
import math
|
| 6 |
|
| 7 |
|
| 8 |
def _normalize_text(text: str) -> str:
|
|
|
|
| 35 |
return 0.0
|
| 36 |
|
| 37 |
|
| 38 |
+
def grade_smart_move(
|
| 39 |
+
predicted: str,
|
| 40 |
+
ground_truth: str,
|
| 41 |
+
label_values: list[str],
|
| 42 |
+
) -> float:
|
| 43 |
+
"""
|
| 44 |
+
Smarter reward for price movement tasks:
|
| 45 |
+
- 40%: Directional Accuracy (Did you get the sign right?)
|
| 46 |
+
- 60%: Numerical Proximity (How close is the percentage?)
|
| 47 |
+
"""
|
| 48 |
+
try:
|
| 49 |
+
# 0. Parsing
|
| 50 |
+
actual_move = float(ground_truth)
|
| 51 |
+
|
| 52 |
+
# Try to parse as JSON first
|
| 53 |
+
predicted_percent = 0.0
|
| 54 |
+
predicted_label = predicted
|
| 55 |
+
if predicted.strip().startswith("{"):
|
| 56 |
+
data = json.loads(predicted)
|
| 57 |
+
predicted_percent = float(data.get("percentage_move", 0.0))
|
| 58 |
+
predicted_label = str(data.get("label", "neutral"))
|
| 59 |
+
else:
|
| 60 |
+
# Fallback: try to extract a float from the string if it's not JSON
|
| 61 |
+
try:
|
| 62 |
+
predicted_percent = float(predicted)
|
| 63 |
+
except ValueError:
|
| 64 |
+
predicted_percent = 0.0
|
| 65 |
+
|
| 66 |
+
# 1. Directional Accuracy (40%)
|
| 67 |
+
# sign(x) is 1 for positive, -1 for negative, 0 for zero
|
| 68 |
+
actual_sign = 0 if abs(actual_move) < 1e-4 else (1 if actual_move > 0 else -1)
|
| 69 |
+
# For simplicity, we compare signs of the numeric percentage if available
|
| 70 |
+
predicted_sign = 0 if abs(predicted_percent) < 1e-4 else (1 if predicted_percent > 0 else -1)
|
| 71 |
+
|
| 72 |
+
directional_reward = 1.0 if actual_sign == predicted_sign else (0.5 if actual_sign == 0 or predicted_sign == 0 else 0.0)
|
| 73 |
+
|
| 74 |
+
# 2. Numerical Proximity (60%)
|
| 75 |
+
# Using exponential decay: exp(-k * error)
|
| 76 |
+
# Scale k: error of 10% (0.1) results in exp(-1.0) ~ 0.36
|
| 77 |
+
k = 10.0
|
| 78 |
+
abs_error = abs(predicted_percent - actual_move)
|
| 79 |
+
numerical_reward = math.exp(-k * abs_error)
|
| 80 |
+
|
| 81 |
+
# 3. Weighted Combination
|
| 82 |
+
return 0.4 * directional_reward + 0.6 * numerical_reward
|
| 83 |
+
|
| 84 |
+
except (json.JSONDecodeError, ValueError, TypeError):
|
| 85 |
+
return 0.0
|
| 86 |
+
|
| 87 |
+
|
| 88 |
def grade_exact(
|
| 89 |
predicted: str,
|
| 90 |
ground_truth: str,
|
tasks/registry.py
CHANGED
|
@@ -14,7 +14,7 @@ from __future__ import annotations
|
|
| 14 |
from collections.abc import Callable
|
| 15 |
from typing import Final
|
| 16 |
|
| 17 |
-
from . import next_quarter_move, sentiment_label
|
| 18 |
from .loader import load_task_subpackage
|
| 19 |
from .types import TaskSpec
|
| 20 |
|
|
@@ -34,8 +34,10 @@ _TASK_ENTRIES: list[tuple[TaskSpec, GradingFn]] = [
|
|
| 34 |
(_pkg_1_day_move.SPEC, _pkg_1_day_move.grade),
|
| 35 |
(_pkg_30_day_move.SPEC, _pkg_30_day_move.grade),
|
| 36 |
(next_quarter_move.SPEC, next_quarter_move.grade),
|
|
|
|
| 37 |
]
|
| 38 |
|
|
|
|
| 39 |
TASKS: dict[str, TaskSpec] = {spec["task_id"]: spec for spec, _ in _TASK_ENTRIES}
|
| 40 |
GRADERS: dict[str, GradingFn] = {spec["task_id"]: fn for spec, fn in _TASK_ENTRIES}
|
| 41 |
|
|
|
|
| 14 |
from collections.abc import Callable
|
| 15 |
from typing import Final
|
| 16 |
|
| 17 |
+
from . import get_figures, next_quarter_move, sentiment_label
|
| 18 |
from .loader import load_task_subpackage
|
| 19 |
from .types import TaskSpec
|
| 20 |
|
|
|
|
| 34 |
(_pkg_1_day_move.SPEC, _pkg_1_day_move.grade),
|
| 35 |
(_pkg_30_day_move.SPEC, _pkg_30_day_move.grade),
|
| 36 |
(next_quarter_move.SPEC, next_quarter_move.grade),
|
| 37 |
+
(get_figures.SPEC, get_figures.grade),
|
| 38 |
]
|
| 39 |
|
| 40 |
+
|
| 41 |
TASKS: dict[str, TaskSpec] = {spec["task_id"]: spec for spec, _ in _TASK_ENTRIES}
|
| 42 |
GRADERS: dict[str, GradingFn] = {spec["task_id"]: fn for spec, fn in _TASK_ENTRIES}
|
| 43 |
|