diff --git a/README.md b/README.md
index ed9aeafb5aba43ed47d8ca57947f5c7aeabe280c..c5765dd40e710a0bad1d5c8665188f5d41b7b041 100644
--- a/README.md
+++ b/README.md
@@ -14,648 +14,3 @@ pinned: false
**Description:** End-to-end financial intelligence pipeline for Net-of-Tax Alpha decisions using causal chain analysis, web scraping, and data aggregation.
---
-
-## Table of Contents
-
-1. [Overview](#overview)
-2. [Authentication](#authentication)
-3. [API Endpoints](#api-endpoints)
- - [Server Endpoints](#server-endpoints)
- - [V2 API Endpoints](#v2-api-endpoints)
-4. [Core Functions](#core-functions)
-5. [Supported Tickers](#supported-tickers)
-6. [Error Handling](#error-handling)
-7. [Examples](#examples)
-
----
-
-## Overview
-
-The Noisy Boy API is a FastAPI-based financial intelligence system that:
-
-- **Fetches multi-source financial data** from various Indian market sources (BSE, NSE, RBI, SEBI, etc.)
-- **Analyzes causal chains** between market events and price movements using BERT and embedding models
-- **Generates aggregated text** from diverse data sources for each ticker
-- **Builds knowledge graphs** with semantic and evidence layers
-- **Streams real-time SSE responses** for long-running analysis tasks
-
-### Technology Stack
-
-- **Framework:** FastAPI with async/await support
-- **Database:** SQLite (local) + Supabase (cloud)
-- **Models:**
- - BERT-CAUSE-EFFECT for causal relationship extraction
- - Fireworks AI embeddings for similarity scoring
- - OpenAI LLM for graph summarization
- - TinyFish AI for web automation
-- **Data Sources:** 50+ fetchers covering macro, regulatory, and company-specific data
-
----
-
-## Authentication
-
-### API Key Authentication
-
-All endpoints (except `/api/source-refs`) require an `X-API-Key` header:
-
-```bash
-curl -H "X-API-Key: your-secret-token" https://api.example.com/api/endpoint
-curl -H "X-API-Key: your-api-key" http://localhost:8000/v2/api/endpoint
-```
-
-**Environment Variable**: `API_KEY` (default: `secret-token`)
-
----
-
-## Base URL
-
-```
-http://localhost:8000/v2
-```
-
----
-
-## API Endpoints
-
-### Scraper API
-
-#### 1. GET `/api/source-refs`
-Returns all fetcher-to-source-URL mappings and reference links.
-
-**Authentication**: None required
-
-**Response**:
-```json
-{
- "fetcher_refs": {
- "repo_rate": "https://rbi.org.in/...",
- "bse": "https://bseindia.com/...",
- ...
- },
- "cat_refs": {
- "macro": {...},
- "corporate": {...}
- }
-}
-```
-
----
-
-#### 2. POST `/api/run`
-Streams real-time scraping results via Server-Sent Events (SSE) from TinyFish API with local caching.
-
-**Authentication**: Required (`X-API-Key`)
-
-**Request Body**:
-```json
-{
- "url": "https://example.com/page",
- "goal": "Extract financial data",
- "ticker": "HDFCBANK",
- "stealth": false
-}
-```
-
-**Query Parameters**:
-- `url` (required): Target URL to scrape
-- `goal` (required): Scraping objective description
-- `ticker` (optional): Stock ticker for caching purposes
-- `stealth` (optional, bool): Enable stealth mode
-
-**Response** (SSE):
-```
-data: {"type": "STARTED", "run_id": "run_123"}
-data: {"type": "PROGRESS", "purpose": "Fetching data..."}
-data: {"type": "COMPLETE", "status": "COMPLETED", "result_json": {...}}
-```
-
-**Caching**: Results are cached per ticker per day via SQLite.
-
----
-
-#### 3. GET `/api/ticker-data/{ticker}`
-Fetch all cached data for a ticker across all sources (Supabase).
-
-**Authentication**: Required (`X-API-Key`)
-
-**Path Parameters**:
-- `ticker`: Stock ticker symbol (e.g., `HDFCBANK`)
-
-**Response**:
-```json
-{
- "ticker": "HDFCBANK",
- "status": "success",
- "data": {
- "repo_rate": {
- "data": {...},
- "fetched_at": "2024-01-15T10:30:00Z"
- },
- "bse": {
- "data": [...],
- "fetched_at": "2024-01-15T10:30:00Z"
- },
- "aggregated_text": {
- "aggregated_text": "...",
- "fetched_at": "2024-01-15T10:30:00Z"
- },
- "causal_chain": {
- "nodes": [...],
- "links": [...],
- "all_chains": [...],
- "biggest_chain": [...]
- }
- }
-}
-```
-
----
-
-### Causal Chain API
-
-#### 4. GET `/api/generate-causal-chain-stream`
-Streams causal chain generation via SSE (two phases: nodes, then edges).
-
-**Authentication**: Required (`X-API-Key`)
-
-**Query Parameters**:
-- `ticker` (required): Stock ticker symbol
-
-**Response** (SSE - Streaming):
-```
-data: {"type": "status", "message": "Generating text for ticker..."}
-data: {"type": "nodes", "nodes": [{"id": "RBI Rate Hike", "label": "RBI Rate Hike"}, ...]}
-data: {"type": "edges", "links": [{"source": "Rate Hike", "target": "Inflation", "score": 0.89}, ...]}
-data: {"type": "done", "all_chains": [[...], [...]], "biggest_chain": [...]}
-```
-
-**Phases**:
-1. **Status**: Initial status message
-2. **Nodes**: Unique causal entities extracted from BERT model
-3. **Edges**: Connections between nodes from embedding similarity
-4. **Done**: Final chains and biggest causal path
-
----
-
-### Ontology API
-
-#### 5. POST `/api/generate-ontology`
-Generate financial ontology from input text.
-
-**Authentication**: Required (via implicit call)
-
-**Request Body**:
-```json
-{
- "ticker": "RELIANCE",
- "text": "Recent crude oil rally drives RELIANCE revenue..."
-}
-```
-
-**Response**:
-```json
-{
- "status": "success",
- "ontology": {
- "entities": ["crude", "revenue", "RELIANCE"],
- "relations": [{"type": "affects", "from": "crude", "to": "RELIANCE"}]
- }
-}
-```
-
----
-
-#### 6. POST `/api/extract-entities`
-Extract named entities using the provided ontology.
-
-**Authentication**: Required (via implicit call)
-
-**Request Body**:
-```json
-{
- "ticker": "ITC",
- "text": "Coal prices surge amid monsoon fears...",
- "ontology": {
- "entities": ["coal", "monsoon", "price"],
- "relations": [...]
- }
-}
-```
-
-**Response**:
-```json
-{
- "status": "success",
- "entities": [
- {"entity": "coal", "type": "commodity", "confidence": 0.92},
- {"entity": "monsoon", "type": "weather_event", "confidence": 0.88}
- ]
-}
-```
-
----
-
-#### 7. POST `/api/build-knowledge-graph`
-Build complete knowledge graph with nodes, edges, and causal chains.
-
-**Authentication**: Required (via implicit call)
-
-**Request Body**:
-```json
-{
- "graph_id": "graph_001",
- "ticker": "BHEL",
- "text": "Full financial text for analysis...",
- "financial_results": {...},
- "forensic_results": {...},
- "tech_results": {...},
- "cached_data": {...}
-}
-```
-
-**Response**:
-```json
-{
- "meta": {
- "ticker": "BHEL",
- "exchange": "NSE",
- "generated_at": "2024-01-15T12:45:00Z",
- "status": "success",
- "chain_count": 5,
- "node_count": 23,
- "edge_count": 45
- },
- "summary": {
- "narrative": "Coal shortage → Power demand surge → BHEL tariff opportunity",
- "net_sentiment_for_ticker": "bullish",
- "ticker_relevance_score": 0.87,
- "macro_regimes_active": ["coal-supply-shock", "power-demand-rally"],
- "top_causal_nodes": ["Coal shortage", "Power demand", "BHEL capacity"],
- "ria_alert": {
- "level": "medium",
- "reason": "Monitor coal supply for sustained impact on tariff structure."
- }
- },
- "biggest_chain": ["Coal shortage", "Power demand", "BHEL contract wins", "Revenue growth"],
- "all_chains": [[...], [...], ...],
- "nodes": [...],
- "edges": [...]
-}
-```
-
----
-
-### Server API (Streaming)
-
-#### 8. GET `/api/generate-causal-chain-stream` (Alternative)
-Same as endpoint #4 but with optional caching from server-side TinyFish integration.
-
----
-
-## Data Fetchers
-
-The system supports multiple data fetchers for different tickers. Each fetcher aggregates specific financial signals:
-
-### Supported Fetchers
-
-| Fetcher | Description | Tickers |
-|---------|-------------|---------|
-| `repo_rate` | RBI Repo Rate | HDFCBANK, TCS, PAYTM, TMPV |
-| `bse` | BSE Announcements | All supported tickers |
-| `fii_dii` | FII/DII Flows | HDFCBANK |
-| `brent` | Brent Crude Oil Price | HDFCBANK, RELIANCE, TMPV |
-| `news` | Financial News Articles | All supported tickers |
-| `npp` | National Power Portal Data | IEX, BHEL |
-| `imd_monsoon` | IMD Monsoon Status | RELIANCE, IEX, BHEL, ETERNAL, ULTRACEMCO |
-| `coal` | Coal Price Index | IEX, BHEL, ITC, ULTRACEMCO |
-| `bhel_tenders` | BHEL Active Tenders | BHEL |
-| `agmarknet` | Agricultural Prices | ETERNAL, ITC |
-| `weather` | Weather in Key Cities | ETERNAL |
-| `h1b` | H1B Visa Filings | TCS |
-| `us_fed` | US Fed Interest Rate | TCS |
-| `us_pmi` | US Services PMI | TCS |
-| `trai_reports` | TRAI Telecom Reports | RELIANCE |
-| `npci` | NPCI UPI Statistics | ETERNAL, IRCTC, PAYTM, MAPMYINDIA |
-| `dgca_traffic` | DGCA Air Traffic Data | IRCTC |
-| `tourist_arrivals` | Foreign Tourist Arrivals | IRCTC |
-| `india_cpi` | India CPI Inflation | ITC |
-| `food_cpi` | Food CPI Inflation | ETERNAL |
-| `labour` | Labour Ministry Releases | ETERNAL |
-| `sebi_orders` | SEBI Orders | PAYTM |
-| `mca_filings` | MCA Corporate Filings | PAYTM |
-| `cma_capacity` | Cement Capacity | ULTRACEMCO |
-| `datareportal` | Digital India Stats | MAPMYINDIA |
-| `pib_highways` | PIB Highway Announcements | ULTRACEMCO, MAPMYINDIA |
-| `pib_vb` | PIB Vande Bharat Updates | BHEL, IRCTC |
-| `nse_bulk_deals` | NSE Bulk Deals | RELIANCE |
-| `cci_orders` | CCI Antitrust Orders | RELIANCE |
-| `erc_orders` | ERC Tariff Orders | IEX |
-| `saubhagya` | Saubhagya Electrification | IEX |
-| `ppac` | India Basket Crude | RELIANCE, ULTRACEMCO |
-
-### Fetcher Response Format
-
-Each fetcher returns formatted text with relevant financial signals:
-
-```
-Brent crude oil price is $85.50 showing a (upward trend, with 5.2% change 30-day).
-RBI Repo Rate is 6.5%.
-FII/DII flows for 2024-01-15: FII net is 450 Cr with net (buy action, holding a positive view).
-```
-
----
-
-## Core Functions
-
-### CausalChain Class (`app/services/causal_chains.py`)
-
-#### `__init__(chunks, fireworks_api_key, fireworks_model)`
-Initialize causal chain processor.
-
-**Parameters**:
-- `chunks` (list): Text chunks to analyze
-- `fireworks_api_key` (str): Fireworks AI API key
-- `fireworks_model` (str): Embedding model name (default: `qwen3-embedding-8b`)
-
----
-
-#### `create_effects(batch_size=16)`
-Extract triggers and effects using BERT-CAUSE-EFFECT model.
-
-**Process**:
-1. Sends text chunks to causal model API
-2. Parses model responses for event triggers and descriptions
-3. Stores triggers and effects
-
-**Retries**: Up to 5 attempts with exponential backoff for failed requests
-
----
-
-#### `create_connections(batch_size=16, chain_threshold=0.85)`
-Build causal connections using embeddings with grounding.
-
-**Process**:
-1. Encodes triggers and effects via Fireworks embeddings
-2. Computes cosine similarity between effects and triggers
-3. Grounds connections by shared entities and keywords
-4. Filters spurious links using entity intersection
-
-**Parameters**:
-- `chain_threshold`: Similarity score threshold (default: 0.85)
-
----
-
-#### `get_all_chains(min_length=2, max_paths=500, time_budget=30)`
-Extract unique causal chains from connections.
-
-**Parameters**:
-- `min_length`: Minimum chain length to keep
-- `max_paths`: Stop after collecting N paths
-- `time_budget`: Wall-clock time limit (seconds)
-
-**Returns**: List of chains, each a list of nodes
-
----
-
-#### `find_biggest_chain(time_budget=20)`
-Find longest causal chain using iterative DFS.
-
-**Returns**: Longest causal path as list of nodes
-
----
-
-#### `to_dict()` / `from_dict()`
-Serialize/deserialize causal chain state.
-
----
-
-### Utility Functions (`util` class)
-
-#### `cos_sim(a, b)`
-Compute cosine similarity matrix between two embedding arrays.
-
-```python
-scores = util.cos_sim(effect_embeddings, trigger_embeddings)
-# scores[i,j] = cosine similarity between effect i and trigger j
-```
-
----
-
-#### `create_chunks(text_input, target_size=700, overlap_sentences=1)`
-Create semantic chunks with overlap for context preservation.
-
-**Parameters**:
-- `target_size`: Target chunk size (characters)
-- `overlap_sentences`: Number of sentences to overlap
-
-**Returns**: List of text chunks ≤ 1000 chars each
-
----
-
-### Fetcher Functions (`scrapper/tiny_fish.py`)
-
-#### `run_fetcher(fetcher_key, ticker_config, persist=False)`
-Execute a single data fetcher.
-
-**Parameters**:
-- `fetcher_key`: Name of fetcher (e.g., `repo_rate`, `bse`)
-- `ticker_config`: Config dict with BSE code, ticker symbol, fetchers list
-- `persist`: Whether to cache results to SQLite
-
-**Returns**: Formatted data dict
-
-**Caching**: Per-ticker, per-day SQLite caching
-
----
-
-### Text Formatting
-
-#### `format_fetcher_text(fetcher_key, data)`
-Format raw fetcher data into human-readable text.
-
-**Example**:
-```python
-text = format_fetcher_text("repo_rate", {"repo_rate_pct": 6.5})
-# Returns: "The current RBI Repo Rate is 6.5%.\n"
-```
-
----
-
-## Examples
-
-### Example 1: Fetch and Cache Aggregated Data
-
-```bash
-curl -X GET "http://localhost:8000/v2/api/ticker-data/RELIANCE" \
- -H "X-API-Key: secret-token"
-```
-
----
-
-### Example 2: Stream Causal Chain Generation (SSE)
-
-```bash
-curl -X GET "http://localhost:8000/v2/api/generate-causal-chain-stream?ticker=HDFCBANK" \
- -H "X-API-Key: secret-token" \
- -N
-```
-
-**Output** (streaming):
-```
-data: {"type":"status","message":"Generating text for ticker..."}
-data: {"type":"nodes","nodes":[{"id":"RBI Rate Hike","label":"RBI Rate Hike"},{"id":"Inflation Risk","label":"Inflation Risk"}]}
-data: {"type":"edges","links":[{"source":"RBI Rate Hike","target":"Inflation Risk","score":0.89}]}
-data: {"type":"done","all_chains":[["RBI Rate Hike","Inflation Risk","Market Correction"]],"biggest_chain":["RBI Rate Hike","Inflation Risk","Market Correction"]}
-```
-
----
-
-### Example 3: Build Knowledge Graph
-
-```bash
-curl -X POST "http://localhost:8000/v2/api/build-knowledge-graph" \
- -H "Content-Type: application/json" \
- -d '{
- "graph_id": "g001",
- "ticker": "TCS",
- "text": "H1B visa approvals surge amid US tech expansion...",
- "financial_results": null,
- "forensic_results": null,
- "tech_results": null,
- "cached_data": null
- }'
-```
-
----
-
-### Example 4: Ontology Generation
-
-```bash
-curl -X POST "http://localhost:8000/v2/api/generate-ontology" \
- -H "Content-Type: application/json" \
- -d '{
- "ticker": "PAYTM",
- "text": "NPCI UPI transaction surge drives fintech growth..."
- }'
-```
-
----
-
-### Example 5: Run Custom Web Scraper
-
-```bash
-curl -X POST "http://localhost:8000/v2/api/run" \
- -H "Content-Type: application/json" \
- -H "X-API-Key: secret-token" \
- -d '{
- "url": "https://example.com/financial-news",
- "goal": "Extract Q4 earnings announcements",
- "ticker": "RELIANCE",
- "stealth": true
- }' \
- -N
-```
-
----
-
-## Error Handling
-
-### Standard Error Responses
-
-#### 400 Bad Request
-```json
-{
- "detail": "url and goal are required"
-}
-```
-
-#### 401 Unauthorized
-```json
-{
- "detail": "Invalid or missing API Key"
-}
-```
-
-#### 500 Internal Server Error
-```json
-{
- "detail": "TINYFISH_API_KEY not found in .env"
-}
-```
-
----
-
-## Supported Tickers
-
-| Ticker | BSE Code | Use Case |
-|--------|----------|----------|
-| HDFCBANK | 500180 | Banking sector, RBI rates |
-| RELIANCE | 500325 | Energy, crude oil exposure |
-| IEX | 540716 | Power sector, capacity data |
-| TCS | 532540 | IT sector, US visa trends |
-| BHEL | 500103 | Power equipment, tender data |
-| ETERNAL | 543320 | Agri-related, food inflation |
-| IRCTC | 542830 | Transportation, tourist flows |
-| ITC | 500875 | Agri-commodities, coal |
-| PAYTM | 543396 | Fintech, NPCI/UPI metrics |
-| ULTRACEMCO | 532538 | Cement, infrastructure |
-| TMPV (TATAMOTORS) | 500570 | Automotive, oil exposure |
-| MAPMYINDIA | 543425 | Digital infrastructure |
-
----
-
-## Response Caching Strategy
-
-- **SQLite Cache**: Per ticker, per day (local fetcher results)
-- **Supabase Cache**: Per ticker, per day (aggregated text, causal chains, ontology)
-- **Cache Invalidation**: Automatic at midnight UTC
-
----
-
-## Performance Notes
-
-- **Causal Chain Generation**: ~30-60 seconds for typical 50-100 chunk inputs
-- **Embedding Computation**: ~5-15 seconds per 100 text chunks
-- **Large Graph Building**: May take 2-5 minutes for complex tickets with >500 nodes
-
----
-
-## Environment Variables
-
-```bash
-API_KEY=secret-token
-FIREWORKS_API_KEY=your-fireworks-key
-CAUSAL_URL=http://localhost:8080/generate_batch
-CAUSAL_API_KEY=optional-causal-model-key
-TINYFISH_API_KEY=your-tinyfish-key
-SUPABASE_URL=https://your-project.supabase.co
-SUPABASE_KEY=your-supabase-key
-DEBUG=false
-```
-
----
-
-## License
-
-Proprietary - ProjectImpulse
-
----
-
-## Support
-
-For issues or questions, contact the ProjectImpulse team.
-```
-
-I've created a comprehensive API README that documents all endpoints, functions, and features in the Noisy_boy repository. The documentation includes:
-
-**Key Sections:**
-1. **Authentication & Base URL** - How to authenticate and where to call endpoints
-2. **8 Major Endpoints** - Scraper, Causal Chain, Ontology, and Server APIs with request/response examples
-3. **30+ Data Fetchers** - Table of all supported data sources for different tickers
-4. **Core Functions** - Detailed documentation of `CausalChain` class, `util` functions, and fetcher operations
-5. **5 Practical Examples** - cURL commands showing how to use each endpoint
-6. **Error Handling** - Standard error response formats
-7. **Performance Notes** - Expected timing for various operations
-8. **Environment Variables** - Complete list of required configs
-
-The documentation follows industry standards with clear formatting, examples, and parameter specifications. You can copy this content into an `API_README.md` file in your repository!
diff --git a/server.py b/server.py
index 55e1e526d56cc72a3571559e79a5eaa2301b1d1e..644479899fedcea74233df72e1d47103a5e83a41 100644
--- a/server.py
+++ b/server.py
@@ -1,7 +1,7 @@
"""
dashboard_server.py
====================
-CUTS+ Causal Terminal — gr.Server entry point.
+Iroha Financial Intelligence — gr.Server entry point.
Architecture
------------
@@ -9,15 +9,9 @@ gr.Server (extends FastAPI)
├── GET / → serves frontend/index.html
├── GET /static/* → serves frontend/{style.css, app.js} (StaticFiles)
│
-├── @server.api run_causal_components → CUTS+ multi-ticker discovery
-├── @server.api run_singular_causal → single-ticker 10-step pipeline
-├── @server.api run_inference → DoFlow / SCM causal query
-├── @server.api run_hierarchy → hierarchical sector causal
+├── @server.api run_inference → DoFlow / SCM causal query (via BACKEND_API)
│
-├── GET /v2/health → health-check
-├── GET /v2/causal/singular-causal/graph/{ticker} → cached adj graph
-├── GET /v2/causal/singular-causal/results/{ticker} → cached inference JSON
-├── POST /v2/causal/doflow-inference → DoFlow query
+├── GET /v2/health → health-check
│
└── All existing /v2/* routers from main.py are included here too
(so this server is a superset of main.py).
@@ -36,9 +30,14 @@ import os
import sys
import json
import logging
+import urllib.error
+import urllib.parse
+import urllib.request
from pathlib import Path
+from gradio import Server
from typing import Any, Dict, List, Optional
-
+from fastapi.staticfiles import StaticFiles
+from fastapi.middleware.cors import CORSMiddleware
from dotenv import load_dotenv
load_dotenv()
@@ -47,7 +46,7 @@ BASE_DIR = Path(__file__).parent.resolve()
if str(BASE_DIR) not in sys.path:
sys.path.insert(0, str(BASE_DIR))
-# Also add the backend directory to sys.path so we can import 'app', 'causal', 'singular_ticker_causal', etc.
+# Also add the backend directory to sys.path so we can import 'app', 'causal', etc.
BACKEND_DIR = (BASE_DIR.parent / "noisy_boy_backend").resolve()
if BACKEND_DIR.exists() and str(BACKEND_DIR) not in sys.path:
sys.path.insert(0, str(BACKEND_DIR))
@@ -65,19 +64,109 @@ logger = logging.getLogger("dashboard-server")
# ─────────────────────────────────────────────────────────────────────────────
# gr.Server
# ─────────────────────────────────────────────────────────────────────────────
-import gradio as gr
-from gradio import Server
-
-from fastapi import Request
-from fastapi.responses import HTMLResponse, FileResponse, JSONResponse
-from fastapi.staticfiles import StaticFiles
-from fastapi.middleware.cors import CORSMiddleware
FRONTEND_DIR = BASE_DIR / "frontend"
INDEX_HTML = FRONTEND_DIR / "index.html"
+# ─────────────────────────────────────────────────────────────────────────────
+# Backend URL
+# ─────────────────────────────────────────────────────────────────────────────
+
+_BACKEND_BASE_URL: str = os.environ.get("BACKEND_API_URL", "http://localhost:7860")
+
+
+# ─────────────────────────────────────────────────────────────────────────────
+# Public API
+# ─────────────────────────────────────────────────────────────────────────────
+
+def run_pipeline(
+ ticker: str = "RELIANCE",
+ threshold: float = 0.5,
+ treatment: Optional[str] = None,
+ outcome: Optional[str] = None,
+ include_pywhyllm: bool = False,
+) -> Dict[str, Any]:
+ """
+ Fetch the validated causal matrix for *ticker* from the backend API.
+
+ Parameters
+ ----------
+ ticker : NSE symbol (e.g. RELIANCE, HDFCBANK)
+ threshold : adjacency threshold for DAG construction
+ treatment : optional treatment node for pywhyllm assumptions
+ outcome : optional outcome node for pywhyllm assumptions
+ include_pywhyllm: request pywhyllm assumption report from backend
+
+ Returns
+ -------
+ dict with keys:
+ nodes, adj_matrix, dag_adj, equations, data_level,
+ topological_order, nodes_graph, links_graph
+ Raises RuntimeError if the backend cannot be reached or returns an error.
+ """
+ params: dict = {"threshold": threshold}
+ if treatment:
+ params["treatment"] = treatment
+ if outcome:
+ params["outcome"] = outcome
+ if include_pywhyllm:
+ params["include_pywhyllm"] = "true"
+
+ qs = urllib.parse.urlencode(params)
+ url = f"{_BACKEND_BASE_URL}/v2/api/singular-causal/causal-matrix/{ticker.upper()}?{qs}"
+ logger.info("run_pipeline: fetching %s", url)
+
+ try:
+ with urllib.request.urlopen(url, timeout=60) as resp:
+ raw = resp.read()
+ except urllib.error.URLError as exc:
+ raise RuntimeError(
+ f"Could not reach backend at {_BACKEND_BASE_URL}. "
+ f"Ensure noisy_boy_backend is running. Original error: {exc}"
+ ) from exc
+
+ payload = json.loads(raw)
+
+ status = payload.get("status")
+ if status == "not_found":
+ raise RuntimeError(
+ payload.get(
+ "detail",
+ f"No cached pipeline data for {ticker} on backend. "
+ "Run the singular-causal pipeline on the backend first.",
+ )
+ )
+ if status not in ("success", None, "ok"):
+ raise RuntimeError(
+ f"Backend returned unexpected status '{status}' for {ticker}. "
+ f"Payload: {payload}"
+ )
+
+ # Build frontend-friendly graph representation
+ nodes: List[str] = payload.get("nodes", [])
+ adj_matrix = payload.get("adj_matrix", [])
+ dag_adj = payload.get("dag_adj", [])
+
+ nodes_graph = [{"id": n, "label": n} for n in nodes]
+ links_graph = []
+ for i, src in enumerate(nodes):
+ for j, dst in enumerate(nodes):
+ if i != j:
+ try:
+ score = float(adj_matrix[i][j])
+ except (IndexError, TypeError, ValueError):
+ score = 0.0
+ if score >= threshold:
+ links_graph.append({"source": src, "target": dst, "score": round(score, 4)})
+
+ return {
+ **payload,
+ "nodes_graph": nodes_graph,
+ "links_graph": links_graph,
+ }
+
server = Server(
- title="CUTS+ Causal Terminal",
+ title="Iroha Causal Terminal",
description=(
"Iroha Financial Intelligence — real-time causal probability matrix, "
"HHKD decomposition, DoFlow inference and sector hierarchy over NIFTY50."
@@ -107,7 +196,7 @@ server.mount(
@server.get("/", response_class=HTMLResponse, include_in_schema=False)
async def serve_index():
- """Serve the CUTS+ Causal Terminal SPA."""
+ """Serve the Iroha Causal Terminal SPA."""
if not INDEX_HTML.exists():
return HTMLResponse("
Frontend not found. Run from backend/
", status_code=500)
return HTMLResponse(INDEX_HTML.read_text(encoding="utf-8"))
@@ -123,78 +212,6 @@ async def health():
return {"status": "ok", "version": "2.0.0"}
-# ─────────────────────────────────────────────────────────────────────────────
-# Existing v2 routers (same as main.py)
-# ─────────────────────────────────────────────────────────────────────────────
-
-try:
- from app.api.scraper import router as scraper_router
- server.include_router(scraper_router, prefix="/v2")
- logger.info("✓ scraper router mounted")
-except Exception as e:
- logger.warning(f"scraper router skipped: {e}")
-
-try:
- from app.api.causal import router as causal_router
- server.include_router(causal_router, prefix="/v2")
- logger.info("✓ causal router mounted")
-except Exception as e:
- logger.warning(f"causal router skipped: {e}")
-
-try:
- from app.api.causal_pipeline import router as causal_pipeline_router
- server.include_router(causal_pipeline_router, prefix="/v2")
- logger.info("✓ causal_pipeline router mounted")
-except Exception as e:
- logger.warning(f"causal_pipeline router skipped: {e}")
-
-try:
- from app.api.screening import router as screening_router
- server.include_router(screening_router, prefix="/v2")
- logger.info("✓ screening router mounted")
-except Exception as e:
- logger.warning(f"screening router skipped: {e}")
-
-try:
- from app.api.ontology import router as ontology_router
- server.include_router(ontology_router, prefix="/v2")
- logger.info("✓ ontology router mounted")
-except Exception as e:
- logger.warning(f"ontology router skipped: {e}")
-
-try:
- # Temporarily remove frontend dir from sys.path to avoid shadowing 'server' package with our 'server.py' script
- removed_empty = False
- if "" in sys.path:
- sys.path.remove("")
- removed_empty = True
- if str(BASE_DIR) in sys.path:
- sys.path.remove(str(BASE_DIR))
-
- from server.model_router import router as server_model_router
- server.include_router(server_model_router, prefix="/v2")
- logger.info("✓ model_router mounted")
-
- # Restore sys.path
- sys.path.insert(0, str(BASE_DIR))
- if removed_empty:
- sys.path.insert(0, "")
-except Exception as e:
- logger.warning(f"model_router skipped: {e}")
- # Ensure sys.path is restored even on failure
- if str(BASE_DIR) not in sys.path:
- sys.path.insert(0, str(BASE_DIR))
- if 'removed_empty' in locals() and removed_empty and "" not in sys.path:
- sys.path.insert(0, "")
-
-# try:
-# from backtest.router import router as backtest_router
-# server.include_router(backtest_router, prefix="/v2/backtest")
-# logger.info("✓ backtest router mounted")
-# except Exception as e:
-# logger.warning(f"backtest router skipped: {e}")
-
-
# ─────────────────────────────────────────────────────────────────────────────
# gr.Server API endpoints (Gradio-backed — queue + SSE streaming)
# These are reachable via the Gradio JS Client as well as plain fetch().
@@ -207,10 +224,6 @@ except Exception as e:
# Override via BACKEND_API_URL env var if running a separate backend on 8000.
_BACKEND_BASE_URL = os.environ.get("BACKEND_API_URL", "http://localhost:7860")
-# Kept for backward-compat with run_singular_causal (which still imports backend modules
-# via sys.path when both repos are co-located). Not used in run_inference anymore.
-_SINGULAR_DEBUG_DIR = str(BASE_DIR / "singular_ticker_causal" / "debug_data")
-
def _fetch_causal_matrix(
ticker: str,
@@ -254,7 +267,10 @@ def _fetch_causal_matrix(
raw = resp.read()
data = json.loads(raw)
if data.get("status") not in ("success", None):
- logger.warning("_fetch_causal_matrix: backend returned status=%s for URL %s. Payload: %s", data.get("status"), url, data)
+ logger.warning(
+ "_fetch_causal_matrix: backend returned status=%s for URL %s. Payload: %s",
+ data.get("status"), url, data,
+ )
return None
return data
except Exception as exc:
@@ -281,147 +297,7 @@ def _safe_json(obj: Any) -> Any:
return obj
-def _adj_to_graph(adj_matrix, symbols: List[str], threshold: float = 0.5):
- """Convert adjacency matrix → {nodes, links} for the frontend."""
- try:
- import numpy as np
- arr = np.array(adj_matrix)
- except Exception:
- arr = [[float(v) for v in row] for row in adj_matrix]
-
- nodes = [{"id": s, "label": s} for s in symbols]
- links = []
- n = len(symbols)
- for i in range(n):
- for j in range(n):
- try:
- v = float(arr[i][j])
- except Exception:
- continue
- if i != j and v >= threshold:
- links.append({"source": symbols[i], "target": symbols[j], "score": round(v, 4)})
- return nodes, links
-
-
-# ── API 1: CUTS+ multi-ticker causal discovery ────────────────────────────
-
-@server.api(name="run_causal_components", description="Run CUTS+ on NIFTY50 multi-ticker technical features")
-def run_causal_components(
- symbols: Optional[List[str]] = None,
- use_actual: bool = False,
-) -> Dict[str, Any]:
- """
- Trigger the CUTS+ multi-ticker causal discovery pipeline.
-
- Parameters
- ----------
- symbols : list of NSE ticker strings (default: synthetic RELIANCE/TCS pair)
- use_actual : whether to download real OHLCV data from yfinance
-
- Returns
- -------
- JSON with adjacency_matrix, nodes, links, density, symbols
- """
- try:
- from causal.test_causal_flow import (
- generate_synthetic_data,
- load_actual_data,
- NIFTY50_SYMBOLS,
- )
- from causal.services.feature_engineer import FeatureEngineer
- from causal.services.cuts_tensor_builder import CutsTensorBuilder
- from causal.cuts_plus.cuts_plus import main as cuts_plus_main
- from causal.cuts_plus.utils.logger import MyLogger
- from omegaconf import OmegaConf
-
- if use_actual:
- syms = symbols or NIFTY50_SYMBOLS
- data = load_actual_data(syms)
- else:
- syms = symbols or ["RELIANCE", "TCS"]
- data = generate_synthetic_data()
-
- fe = FeatureEngineer()
- ctb = CutsTensorBuilder()
- tech_data, mask, ordered_syms, *_ = ctb.build(
- historical_data=data, symbols=syms, feature_engineer=fe
- )
-
- log_dir = str(BASE_DIR / "causal" / "dash_logs")
- os.makedirs(log_dir, exist_ok=True)
- log = MyLogger(log_dir=log_dir, stdout=False, stderr=False, tensorboard=False)
-
- cfg = OmegaConf.create({
- "data_dim": tech_data.shape[-1],
- "total_epoch": 30,
- "ticker_list": ordered_syms,
- "causal_thres": "value_0.5",
- })
-
- adj = cuts_plus_main(data=tech_data, mask=mask, true_cm=None, opt=cfg, log=log)
- adj_list = _safe_json(adj)
- nodes, links = _adj_to_graph(adj_list, ordered_syms, threshold=0.5)
- n = len(ordered_syms)
- edges = sum(1 for i in range(n) for j in range(n) if i != j and adj_list[i][j] >= 0.5)
-
- return {
- "status": "ok",
- "symbols": ordered_syms,
- "adjacency_matrix": adj_list,
- "nodes": nodes,
- "links": links,
- "density": round(edges / max(n * (n - 1), 1), 4),
- "n_edges": edges,
- }
- except Exception as exc:
- logger.exception("run_causal_components failed")
- return {"status": "error", "detail": str(exc)}
-
-
-# ── API 2: Single-ticker fundamental causal pipeline ─────────────────────
-
-@server.api(
- name="run_singular_causal",
- description="Run the full 10-step single-ticker fundamental causal pipeline",
- concurrency_limit=2,
-)
-def run_singular_causal(ticker: str = "RELIANCE") -> Dict[str, Any]:
- """
- Execute the 10-step singular ticker pipeline:
- DuPont prior → feature engineering → CUTS+ → SCM → CausalQueryEngine.
-
- Parameters
- ----------
- ticker : NSE symbol (e.g. RELIANCE, HDFCBANK)
-
- Returns
- -------
- JSON with adj_matrix, nodes, links, inference_summary
- """
- try:
- from singular_ticker_causal.test_single_ticker_causal_flow import run_pipeline
- result = run_pipeline(ticker=ticker.upper())
- return _safe_json({"status": "ok", "ticker": ticker.upper(), **result})
- except Exception as exc:
- logger.exception("run_singular_causal failed")
- return {"status": "error", "detail": str(exc)}
-
-
-# ── API 3: Causal inference (assert / intervene / counterfactual) ─────────
-#
-# Architecture:
-# 1. Fetch the VALIDATED causal matrix from noisy_boy_backend via HTTP.
-# The backend has already run CUTS+ learning + pywhyllm + DoWhy validation.
-# 2. Reconstruct the fitted SCM locally from that payload (no re-learning).
-# 3. Use pywhyllm to gather structural guidance (confounders, backdoor sets,
-# SCM mechanism hints) at each of the three causal layers.
-# 4. Feed that guidance + the fitted SCM data into DoWhy / DoWhy-GCM to
-# compute the actual numerical estimates — the LLM never touches the numbers.
-
-
-def _resolve_value(
- value: float, value_type: str, current: float
-) -> float:
+def _resolve_value(value: float, value_type: str, current: float) -> float:
"""Convert a user-supplied value + value_type to the absolute node value."""
vt = value_type.strip().lower()
if vt == "absolute":
@@ -434,81 +310,168 @@ def _resolve_value(
return value
-def _rebuild_scm_from_payload(payload: dict):
- """
- Reconstruct a fitted StructuralCausalModel from the causal-matrix payload.
-
- The backend has already:
- - run CUTS+ to learn the adjacency matrix
- - fit the structural equations (coefficients, intercepts, residual_std)
- - validated the graph with pywhyllm + DoWhy refutation
+# ─────────────────────────────────────────────────────────────────────────────
+# Pure-numpy inference helpers (no local causal training imports)
+# These functions work entirely from the payload returned by the backend API.
+# ─────────────────────────────────────────────────────────────────────────────
- We re-hydrate a StructuralCausalModel object from that payload so that the
- frontend inference code can call engine.assert_edge / intervene / counterfactual
- without re-running any learning.
+def _build_dag_from_payload(payload: dict):
+ """
+ Return a numpy bool DAG adjacency matrix and list of node names
+ from the backend causal-matrix payload.
"""
import numpy as np
- from singular_ticker_causal.causal_inference.causal_model import (
- StructuralCausalModel, StructuralEquation
- )
-
nodes = payload["nodes"]
- adj_matrix = np.array(payload["adj_matrix"], dtype=float)
dag_adj = np.array(payload["dag_adj"], dtype=bool)
- data_level = np.array(payload["data_level"], dtype=float)
+ adj_matrix = np.array(payload["adj_matrix"], dtype=float)
+ return nodes, dag_adj, adj_matrix
+
+
+def _propagate_intervention(
+ nodes: list,
+ dag_adj,
+ equations: dict,
+ data_level,
+ topological_order: list,
+ treatment: str,
+ abs_value: float,
+ targets: list,
+ horizon: int = 5,
+):
+ """
+ Propagate a hard intervention (do(treatment=abs_value)) through the
+ structural equations for `horizon` steps, returning ATE per target node.
+ Uses only numpy — no local causal model imports.
+ """
+ import numpy as np
+
+ node_to_idx = {n: i for i, n in enumerate(nodes)}
n = len(nodes)
T = data_level.shape[0]
- # Build a minimal (T, N, 1) data_tech tensor so StructuralCausalModel.__post_init__
- # can call _extract_level_data without error. The level data IS data_level.
- data_tech = data_level[:, :, np.newaxis] # shape (T, N, 1)
- adjacency_mask = (adj_matrix > 0).astype(float) # use adj as mask
-
- scm = StructuralCausalModel(
- nodes=nodes,
- adj=adj_matrix,
- adjacency_mask=adjacency_mask,
- data_tech=data_tech,
- threshold=payload.get("threshold", 0.5),
- lag=1,
- )
+ # Start from the last observed time step
+ state = data_level[-1].copy().astype(float)
- # Override dag_adj with the backend's thresholded version
- scm.dag_adj = dag_adj
-
- # Restore topological order
- topo_names = payload.get("topological_order", nodes)
- scm.topological_indices = [nodes.index(n) for n in topo_names if n in nodes]
-
- # Re-hydrate structural equations from the backend payload
- equations_raw = payload.get("equations", {})
- scm.equations = {}
- for node, eq_data in equations_raw.items():
- parents = eq_data.get("parents", [])
- parent_indices = eq_data.get("parent_indices", [nodes.index(p) for p in parents])
- scm.equations[node] = StructuralEquation(
- node=node,
- parents=parents,
- parent_indices=parent_indices,
- intercept=float(eq_data.get("intercept", 0.0)),
- coefficients={p: float(v) for p, v in eq_data.get("coefficients", {}).items()},
- residual_mean=float(eq_data.get("residual_mean", 0.0)),
- residual_std=float(eq_data.get("residual_std", 1.0)),
- r_squared=float(eq_data.get("r_squared", 0.0)),
- n_obs=int(eq_data.get("n_obs", T)),
- equation_type=eq_data.get("equation_type", "linear"),
- )
+ # Fix the treatment node
+ t_idx = node_to_idx[treatment]
+ state[t_idx] = abs_value
- return scm
+ ate_per_target: Dict[str, float] = {}
+ baseline = data_level[-1].copy().astype(float)
+
+ for _ in range(horizon):
+ new_state = state.copy()
+ for node_name in topological_order:
+ if node_name == treatment:
+ continue
+ eq = equations.get(node_name)
+ if eq is None:
+ continue
+ parents = eq.get("parents", [])
+ coefficients = eq.get("coefficients", {})
+ intercept = float(eq.get("intercept", 0.0))
+ if not parents:
+ continue
+ val = intercept
+ for p in parents:
+ p_idx = node_to_idx.get(p)
+ if p_idx is not None:
+ val += float(coefficients.get(p, 0.0)) * float(state[p_idx])
+ n_idx = node_to_idx[node_name]
+ new_state[n_idx] = val
+ state = new_state
+
+ for target in targets:
+ t_i = node_to_idx.get(target)
+ if t_i is not None:
+ ate_per_target[target] = float(state[t_i] - baseline[t_i])
+
+ return ate_per_target, state
+
+
+def _abduct_and_predict(
+ nodes: list,
+ dag_adj,
+ equations: dict,
+ data_level,
+ topological_order: list,
+ treatment: str,
+ cf_value: float,
+ target: str,
+ observed_t: int,
+):
+ """
+ Simple SCM abduction for counterfactual:
+ 1. Abduct residuals from the observed time step.
+ 2. Re-run structural equations with treatment fixed to cf_value.
+ 3. Return factual_outcome, cf_outcome, ITE.
+ """
+ import numpy as np
+
+ node_to_idx = {n: i for i, n in enumerate(nodes)}
+ obs = data_level[observed_t].copy().astype(float)
+
+ # Abduct residuals
+ residuals: Dict[str, float] = {}
+ for node_name in topological_order:
+ eq = equations.get(node_name)
+ if eq is None or not eq.get("parents"):
+ residuals[node_name] = 0.0
+ continue
+ parents = eq.get("parents", [])
+ coefficients = eq.get("coefficients", {})
+ intercept = float(eq.get("intercept", 0.0))
+ predicted = intercept
+ for p in parents:
+ p_idx = node_to_idx.get(p)
+ if p_idx is not None:
+ predicted += float(coefficients.get(p, 0.0)) * float(obs[node_to_idx[p]])
+ residuals[node_name] = float(obs[node_to_idx[node_name]]) - predicted
+
+ # Counterfactual: fix treatment, replay equations with abducted noise
+ cf_state = obs.copy()
+ cf_state[node_to_idx[treatment]] = cf_value
+
+ for node_name in topological_order:
+ if node_name == treatment:
+ continue
+ eq = equations.get(node_name)
+ if eq is None or not eq.get("parents"):
+ continue
+ parents = eq.get("parents", [])
+ coefficients = eq.get("coefficients", {})
+ intercept = float(eq.get("intercept", 0.0))
+ predicted = intercept
+ for p in parents:
+ p_idx = node_to_idx.get(p)
+ if p_idx is not None:
+ predicted += float(coefficients.get(p, 0.0)) * float(cf_state[p_idx])
+ n_idx = node_to_idx[node_name]
+ cf_state[n_idx] = predicted + residuals.get(node_name, 0.0)
+
+ factual_outcome = float(obs[node_to_idx[target]])
+ cf_outcome = float(cf_state[node_to_idx[target]])
+ ite = cf_outcome - factual_outcome
+ return factual_outcome, cf_outcome, ite
+
+
+# ── API: Causal inference (assert / intervene / counterfactual) ───────────
+#
+# Architecture:
+# 1. Fetch the VALIDATED causal matrix from noisy_boy_backend via HTTP.
+# The backend has already run CUTS+ learning + pywhyllm + DoWhy validation.
+# 2. Use the payload data (equations, adj, data_level) for inference
+# using pure numpy/pandas — no local causal training imports required.
+# 3. Optionally consult pywhyllm guidance from the backend payload.
@server.api(
name="run_inference",
description=(
- "Run pywhyllm-guided causal inference (association / intervention / counterfactual) "
- "using a validated causal matrix from the backend. "
- "Layers: 1=Association(pandas+DoWhy), 2=Intervention(DoWhy backdoor/IV), "
- "3=Counterfactual(DoWhy GCM with abduction)."
+ "Run causal inference (association / intervention / counterfactual) "
+ "using a validated causal matrix fetched from the backend API. "
+ "Layers: 1=Association(DoWhy backdoor), 2=Intervention(SCM propagation), "
+ "3=Counterfactual(SCM abduction)."
),
concurrency_limit=4,
)
@@ -587,35 +550,38 @@ def run_inference(
"detail": payload.get("detail", f"No cached pipeline data for {ticker}."),
}
- # ── 2. Reconstruct fitted SCM from payload (no re-learning) ───────────
- sys.path.insert(0, str(BASE_DIR)) # ensure singular_ticker_causal is importable
- scm = _rebuild_scm_from_payload(payload)
+ # ── 2. Unpack payload (no local causal training imports) ──────────────
+ nodes, dag_adj, adj_matrix = _build_dag_from_payload(payload)
+ node_to_idx = {n: i for i, n in enumerate(nodes)}
+ data_level = np.array(payload["data_level"], dtype=float)
+ equations_raw = payload.get("equations", {})
+ topo_order = payload.get("topological_order", nodes)
- from singular_ticker_causal.causal_inference.query_engine import CausalQueryEngine
+ T = data_level.shape[0]
+ df = pd.DataFrame(data_level, columns=nodes)
- # Attach pywhyllm config if requested
- engine = CausalQueryEngine(
- scm,
- pywhyllm_enabled=use_pywhyllm,
- )
+ if treatment not in node_to_idx:
+ return {"status": "error", "detail": f"Unknown treatment node: {treatment}"}
+ if target_node not in node_to_idx:
+ return {"status": "error", "detail": f"Unknown outcome/target node: {target_node}"}
+ if df.shape[0] < 5:
+ return {
+ "status": "error",
+ "detail": f"Insufficient observations ({df.shape[0]}) to run inference.",
+ }
- # ── 3. pywhyllm structural guidance (Layer-aware) ─────────────────────
- # pywhyllm identifies confounders, backdoor sets, and mechanism hints.
- # It NEVER computes the final number — that is DoWhy's job.
- pywhyllm_report: Optional[dict] = payload.get("pywhyllm_report") # pre-fetched if requested
+ # ── 3. pywhyllm structural guidance (from backend payload) ────────────
+ pywhyllm_report: Optional[dict] = payload.get("pywhyllm_report")
adjustment_sets: List[List[str]] = []
- suggested_ivs: List[str] = []
if use_pywhyllm and pywhyllm_report and pywhyllm_report.get("available"):
- # Extract backdoor adjustment candidates suggested by the LLM
raw_backdoor = pywhyllm_report.get("suggested_backdoor_sets") or []
- valid_nodes = set(scm.nodes) - {treatment, target_node}
+ valid_nodes = set(nodes) - {treatment, target_node}
for suggested_set in raw_backdoor:
clean = [n for n in suggested_set if n in valid_nodes]
if clean and clean not in adjustment_sets:
adjustment_sets.append(clean)
- # Also pick up confounder suggestions as a fallback adjustment set
confounders = [
n for n in (pywhyllm_report.get("suggested_confounders") or [])
if n in valid_nodes
@@ -623,36 +589,24 @@ def run_inference(
if confounders and confounders not in adjustment_sets:
adjustment_sets.append(confounders)
- # Instrumental variables (for Layer 2 IV estimation)
- suggested_ivs = [
- n for n in (pywhyllm_report.get("suggested_ivs") or [])
- if n in scm.node_to_idx
- ]
-
- df = pd.DataFrame(scm.data_level, columns=scm.nodes)
- if df.shape[0] < 5:
- return {
- "status": "error",
- "detail": f"Insufficient observations ({df.shape[0]}) to run inference.",
- }
-
result: Dict[str, Any] = {}
# ═══════════════════════════════════════════════════════════════════════
# LAYER 1 — Association: "What does Y look like given X?"
- # pywhyllm role: identify confounders and suggest adjustment variables
- # execution: DoWhy identifies + estimates via backdoor linear regression
+ # Uses DoWhy with the backend-provided DAG, falling back to OLS.
# ═══════════════════════════════════════════════════════════════════════
if mode == "assert":
- if treatment not in scm.node_to_idx:
- return {"status": "error", "detail": f"Unknown treatment node: {treatment}"}
- if target_node not in scm.node_to_idx:
- return {"status": "error", "detail": f"Unknown outcome node: {target_node}"}
-
try:
from dowhy import CausalModel
- graph_dot = engine._build_dowhy_graph()
+ # Build DOT graph string from dag_adj
+ edges = []
+ for si, src in enumerate(nodes):
+ for di, dst in enumerate(nodes):
+ if dag_adj[si, di]:
+ edges.append(f"{src} -> {dst}")
+ graph_dot = "digraph{" + "; ".join(edges) + "}"
+
dowhy_model = CausalModel(
data=df,
treatment=treatment,
@@ -668,18 +622,15 @@ def run_inference(
)
ate = float(estimate.value)
- # Real confidence interval from the linear model's standard error
- # DoWhy stores the sklearn estimator under estimate.estimator
+ # Confidence interval from OLS residuals
se: float = 0.0
try:
- est_obj = estimate.estimator
- # Compute SE from coefficient covariance if available
+ import numpy.linalg as nla
X = df[[c for c in df.columns if c != target_node]].values
y = df[target_node].values
- import numpy.linalg as nla
XtX_inv = nla.pinv(X.T @ X)
resid = y - X @ nla.lstsq(X, y, rcond=None)[0]
- sigma2 = float(np.sum(resid**2) / max(1, len(y) - X.shape[1]))
+ sigma2 = float(np.sum(resid ** 2) / max(1, len(y) - X.shape[1]))
t_idx_local = list(df.columns).index(treatment)
se = float(np.sqrt(max(0.0, sigma2 * XtX_inv[t_idx_local, t_idx_local])))
except Exception:
@@ -689,14 +640,14 @@ def run_inference(
ci_upper = ate + 1.96 * se
prob = min(1.0, abs(ate) / (abs(ate) + se + 1e-9))
- # Ripple effects: downstream nodes reachable from treatment in the DAG
+ # Ripple effects: direct downstream neighbours of treatment
ripple_effects = []
- t_idx_scm = scm.node_to_idx[treatment]
- for j, node in enumerate(scm.nodes):
+ t_idx_g = node_to_idx[treatment]
+ for j, node in enumerate(nodes):
if node == treatment or node == target_node:
continue
- if scm.dag_adj[t_idx_scm, j]:
- edge_score = float(scm.adj[t_idx_scm, j])
+ if dag_adj[t_idx_g, j]:
+ edge_score = float(adj_matrix[t_idx_g, j])
ripple_effects.append({
"ticker": node,
"direction": 1 if ate > 0 else -1,
@@ -714,47 +665,65 @@ def run_inference(
}
except Exception as dowhy_exc:
- # DoWhy not installed or identification failed — fall back to SCM engine
- logger.warning("DoWhy association failed (%s), falling back to SCM", dowhy_exc)
- scm_result = engine.assert_edge(treatment, target_node)
- ci = scm_result.get("ci_95", (0.0, 0.0))
+ # DoWhy not installed or identification failed — fall back to OLS
+ logger.warning("DoWhy association failed (%s), falling back to OLS", dowhy_exc)
+ t_idx_g = node_to_idx[treatment]
+ out_idx = node_to_idx[target_node]
+
+ # Simple OLS: regress target on treatment
+ X = df[[treatment]].values
+ y = df[target_node].values
+ import numpy.linalg as nla
+ coef = nla.lstsq(np.c_[np.ones(len(X)), X], y, rcond=None)[0]
+ ate = float(coef[1])
+ se = abs(ate) * 0.15
+ ci_lower = ate - 1.96 * se
+ ci_upper = ate + 1.96 * se
+
+ ripple_effects = []
+ for j, node in enumerate(nodes):
+ if node == treatment or node == target_node:
+ continue
+ if dag_adj[t_idx_g, j]:
+ ripple_effects.append({
+ "ticker": node,
+ "direction": 1 if ate > 0 else -1,
+ "magnitude": round(float(adj_matrix[t_idx_g, j]) * abs(ate), 4),
+ })
+
result = {
- "ate": scm_result.get("ate", 0.0),
- "ci_lower": ci[0],
- "ci_upper": ci[1],
- "probability": min(1.0, abs(scm_result.get("ate", 0.0))),
- "strategy": scm_result.get("strategy", "scm_fallback"),
- "adjustment_set": sorted(scm_result.get("adjustment_set") or []),
- "ripple_effects": [],
+ "ate": ate,
+ "ci_lower": ci_lower,
+ "ci_upper": ci_upper,
+ "probability": min(1.0, abs(ate) / (abs(ate) + se + 1e-9)),
+ "strategy": "ols_fallback",
+ "adjustment_set": adjustment_sets[0] if adjustment_sets else [],
+ "ripple_effects": ripple_effects,
}
# ═══════════════════════════════════════════════════════════════════════
# LAYER 2 — Intervention: "What will happen to Y if we do X=value?"
- # pywhyllm role: suggest backdoor sets and IV strategy
- # execution: DoWhy identifies + estimates; SCM engine propagates ripples
+ # Propagates through structural equations from the backend payload.
# ═══════════════════════════════════════════════════════════════════════
elif mode == "intervene":
- if treatment not in scm.node_to_idx:
- return {"status": "error", "detail": f"Unknown treatment node: {treatment}"}
- if target_node not in scm.node_to_idx:
- return {"status": "error", "detail": f"Unknown outcome node: {target_node}"}
-
- # Resolve the absolute intervention value
- current_val = float(scm.data_level[-1, scm.node_to_idx[treatment]])
+ current_val = float(data_level[-1, node_to_idx[treatment]])
abs_value = _resolve_value(value, value_type, current_val)
- # ── DoWhy: estimate the causal effect under the intervention ──────
+ # Try DoWhy for ATE estimation first
+ ate = 0.0
+ method_used = "scm_propagation"
try:
from dowhy import CausalModel
- graph_dot = engine._build_dowhy_graph()
-
- # Build a modified dataset where treatment is fixed to abs_value
- df_intervened = df.copy()
- df_intervened[treatment] = abs_value
+ edges = []
+ for si, src in enumerate(nodes):
+ for di, dst in enumerate(nodes):
+ if dag_adj[si, di]:
+ edges.append(f"{src} -> {dst}")
+ graph_dot = "digraph{" + "; ".join(edges) + "}"
dowhy_model = CausalModel(
- data=df, # use original data for identification
+ data=df,
treatment=treatment,
outcome=target_node,
graph=graph_dot,
@@ -762,63 +731,39 @@ def run_inference(
identified_estimand = dowhy_model.identify_effect(
proceed_when_unidentifiable=True
)
-
- # Use IV estimator if pywhyllm suggested one, else backdoor
- if suggested_ivs:
- try:
- estimate = dowhy_model.estimate_effect(
- identified_estimand,
- method_name="iv.instrumental_variable",
- method_params={"iv_instrument_name": suggested_ivs[0]},
- )
- method_used = f"iv.instrumental_variable ({suggested_ivs[0]})"
- except Exception:
- estimate = dowhy_model.estimate_effect(
- identified_estimand,
- method_name="backdoor.linear_regression",
- )
- method_used = "backdoor.linear_regression (IV fallback)"
- else:
- estimate = dowhy_model.estimate_effect(
- identified_estimand,
- method_name="backdoor.linear_regression",
- )
- method_used = "backdoor.linear_regression"
-
- # Scale the ATE by the actual intervention delta
- ate_unit = float(estimate.value) # effect per unit of treatment
+ estimate = dowhy_model.estimate_effect(
+ identified_estimand,
+ method_name="backdoor.linear_regression",
+ )
+ ate_unit = float(estimate.value)
delta = abs_value - current_val
ate = ate_unit * delta
-
- # SE estimation
- se = abs(ate) * 0.12
- ci_lower = ate - 1.96 * se
- ci_upper = ate + 1.96 * se
-
+ method_used = "backdoor.linear_regression"
except Exception as dowhy_exc:
- logger.warning("DoWhy intervention failed (%s), using SCM engine", dowhy_exc)
- method_used = "scm_propagation"
- ate = 0.0
- ci_lower = 0.0
- ci_upper = 0.0
-
- # ── SCM engine: propagate intervention to get ripple effects ──────
- scm_int_result = engine.intervene(
+ logger.warning("DoWhy intervention failed (%s), using SCM propagation", dowhy_exc)
+
+ # SCM propagation for ripple effects (pure numpy, no training imports)
+ ate_per_target, final_state = _propagate_intervention(
+ nodes=nodes,
+ dag_adj=dag_adj,
+ equations=equations_raw,
+ data_level=data_level,
+ topological_order=topo_order,
treatment=treatment,
- value=abs_value,
- targets=[target_node],
+ abs_value=abs_value,
+ targets=[target_node] + [n for n in nodes if n != treatment],
horizon=horizon,
)
- ate_per_target = scm_int_result.get("ate_per_target", {})
- if target_node in ate_per_target and method_used == "scm_propagation":
+ if method_used == "scm_propagation" and target_node in ate_per_target:
ate = float(ate_per_target[target_node])
- ci_lower = ate - abs(ate) * 0.15
- ci_upper = ate + abs(ate) * 0.15
- # Build ripple effects from all downstream SCM targets
+ se = abs(ate) * 0.12
+ ci_lower = ate - 1.96 * se
+ ci_upper = ate + 1.96 * se
+
ripple_effects = []
- for node, delta_val in (scm_int_result.get("ate_per_target") or {}).items():
+ for node, delta_val in ate_per_target.items():
if node == treatment:
continue
ripple_effects.append({
@@ -836,75 +781,51 @@ def run_inference(
"intervention_value": abs_value,
"value_type": value_type,
"horizon": horizon,
- "predicted_values": _safe_json(scm_int_result.get("predicted_values", {})),
"ripple_effects": ripple_effects,
"adjustment_set": adjustment_sets[0] if adjustment_sets else [],
}
# ═══════════════════════════════════════════════════════════════════════
# LAYER 3 — Counterfactual: "What if X had been different in the past?"
- # pywhyllm role: formulate SCM mechanism assignments for GCM
- # execution: DoWhy GCM abducts noise → applies counterfactual → predicts
+ # Uses SCM abduction via pure numpy structural equations.
# ═══════════════════════════════════════════════════════════════════════
elif mode in ("counterfactual", "counter"):
- if treatment not in scm.node_to_idx:
- return {"status": "error", "detail": f"Unknown treatment node: {treatment}"}
- if target_node not in scm.node_to_idx:
- return {"status": "error", "detail": f"Unknown target node: {target_node}"}
-
# Resolve observed timestep
- T = scm.t_steps
t = observed_t if observed_t >= 0 else (T + observed_t)
t = max(0, min(T - 1, t))
# Resolve counterfactual value
- current_val = float(scm.data_level[t, scm.node_to_idx[treatment]])
+ current_val = float(data_level[t, node_to_idx[treatment]])
if cf_value is not None:
abs_cf_value = float(cf_value)
else:
abs_cf_value = _resolve_value(value, value_type, current_val)
- # ── DoWhy GCM counterfactual (primary path) ───────────────────────
+ # Try DoWhy GCM first
gcm_used = False
+ factual_outcome = 0.0
+ cf_outcome_val = 0.0
+ ite = 0.0
+
try:
import dowhy.gcm as gcm_module
import networkx as nx
- # Build directed causal graph from the validated DAG
causal_graph = nx.DiGraph()
- for src_i, src_name in enumerate(scm.nodes):
- for dst_i, dst_name in enumerate(scm.nodes):
- if scm.dag_adj[src_i, dst_i]:
- causal_graph.add_edge(src_name, dst_name)
- for node in scm.nodes:
+ for si, src in enumerate(nodes):
+ for di, dst in enumerate(nodes):
+ if dag_adj[si, di]:
+ causal_graph.add_edge(src, dst)
+ for node in nodes:
if node not in causal_graph.nodes:
causal_graph.add_node(node)
- # pywhyllm guidance: use equation types to assign mechanisms
- # - Nodes with parents get AdditiveNoiseModel (invertible, required for CF)
- # - Root (exogenous) nodes get EmpiricalDistribution
gcm_model = gcm_module.InvertibleStructuralCausalModel(causal_graph)
gcm_module.auto.assign_mechanisms(gcm_model, df)
-
- # Override mechanism types based on pywhyllm's equation suggestions
- # if available, to improve SCM quality
- if pywhyllm_report and pywhyllm_report.get("available"):
- for node in scm.nodes:
- eq_data = (payload.get("equations") or {}).get(node, {})
- if eq_data.get("equation_type") == "exogenous":
- if node in gcm_model.graph.nodes:
- gcm_model.set_causal_mechanism(
- node,
- gcm_module.EmpiricalDistribution()
- )
-
gcm_module.fit(gcm_model, df)
- # The observed data at time t
observed_data = df.iloc[[t]]
-
- # Run counterfactual: fix treatment, abduct noise, predict
- cf_val_fixed = abs_cf_value # capture in closure
+ cf_val_fixed = abs_cf_value
cf_samples = gcm_module.counterfactual_samples(
gcm_model,
{treatment: lambda x, v=cf_val_fixed: np.full(x.shape, v)},
@@ -913,42 +834,32 @@ def run_inference(
)
factual_outcome = float(observed_data[target_node].iloc[0])
- cf_outcome = float(cf_samples[target_node].iloc[0])
- ite = cf_outcome - factual_outcome
+ cf_outcome_val = float(cf_samples[target_node].iloc[0])
+ ite = cf_outcome_val - factual_outcome
gcm_used = True
except Exception as gcm_exc:
logger.warning("DoWhy GCM counterfactual failed (%s), using SCM abduction", gcm_exc)
- gcm_used = False
if not gcm_used:
- # ── Fallback: SCM abduction engine (always available) ─────────
- scm_cf_result = engine.counterfactual(
- observed_t=t,
+ factual_outcome, cf_outcome_val, ite = _abduct_and_predict(
+ nodes=nodes,
+ dag_adj=dag_adj,
+ equations=equations_raw,
+ data_level=data_level,
+ topological_order=topo_order,
treatment=treatment,
cf_value=abs_cf_value,
target=target_node,
- )
- factual_outcome = float(scm_cf_result.get("factual_outcome", 0.0))
- cf_outcome = float(scm_cf_result.get("counterfactual_outcome", 0.0))
- ite = float(scm_cf_result.get("ite", 0.0))
-
- # Shapley contributions — always from SCM engine (numerically exact)
- shapley: Dict[str, float] = {}
- try:
- shapley_result = engine.counterfactual(
observed_t=t,
- treatment=treatment,
- cf_value=abs_cf_value,
- target=target_node,
)
- shapley = shapley_result.get("shapley_contributions", {treatment: ite})
- except Exception:
- shapley = {treatment: ite}
- # SE from residual std of the target equation
- target_eq = scm.equations.get(target_node)
- se = float(target_eq.residual_std) if target_eq else abs(ite) * 0.15
+ # Shapley: single-treatment — just use the ITE directly
+ shapley = {treatment: ite}
+
+ # SE from residual_std of the target equation (from backend payload)
+ target_eq_data = equations_raw.get(target_node, {})
+ se = float(target_eq_data.get("residual_std", abs(ite) * 0.15))
ci_lower = ite - 1.96 * se
ci_upper = ite + 1.96 * se
@@ -956,7 +867,7 @@ def run_inference(
"ate": ite,
"ite": ite,
"factual_outcome": factual_outcome,
- "counterfactual_outcome": cf_outcome,
+ "counterfactual_outcome": cf_outcome_val,
"ci_lower": ci_lower,
"ci_upper": ci_upper,
"probability": min(1.0, abs(ite) / (abs(ite) + se + 1e-9)),
@@ -964,7 +875,7 @@ def run_inference(
"counterfactual_value": abs_cf_value,
"value_type": value_type,
"observed_t": t,
- "shapley_contributions": _safe_json(shapley),
+ "shapley_contributions": shapley,
"ripple_effects": [],
}
@@ -985,36 +896,6 @@ def run_inference(
return {"status": "error", "detail": str(exc)}
-# ── API 4: Hierarchical sector causal ────────────────────────────────────
-
-@server.api(
- name="run_hierarchy",
- description="Run CrossLevelMPNN hierarchical sector causal graph",
- concurrency_limit=1,
-)
-def run_hierarchy(
- symbols: Optional[List[str]] = None,
-) -> Dict[str, Any]:
- """
- Build the micro + macro causal hierarchy graph.
-
- Parameters
- ----------
- symbols : Optional override for the symbol list (defaults to top-4 NIFTY tickers)
-
- Returns
- -------
- JSON with micro_graph, macro_graph, sector_embeddings
- """
- try:
- from causal_hierarchy.test_hierarchical_causal_flow import run_hierarchical_flow
- result = run_hierarchical_flow(symbols=symbols)
- return _safe_json({"status": "ok", **result})
- except Exception as exc:
- logger.exception("run_hierarchy failed")
- return {"status": "error", "detail": str(exc)}
-
-
# ─────────────────────────────────────────────────────────────────────────────
# Entry point
# ─────────────────────────────────────────────────────────────────────────────
@@ -1023,7 +904,7 @@ if __name__ == "__main__":
port = int(os.environ.get("GRADIO_SERVER_PORT", os.environ.get("PORT", "7860")))
host = os.environ.get("GRADIO_SERVER_NAME", "0.0.0.0")
- logger.info(f"Starting CUTS+ Causal Terminal on {host}:{port}")
+ logger.info(f"Starting Iroha Causal Terminal on {host}:{port}")
logger.info(f" → Frontend : http://localhost:{port}/")
logger.info(f" → API docs : http://localhost:{port}/docs")
diff --git a/singular_ticker_causal/.env b/singular_ticker_causal/.env
deleted file mode 100644
index 0691c52dce7e7a857d2f0f745541d612d6e344a0..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/.env
+++ /dev/null
@@ -1,3 +0,0 @@
-FIREWORKS_API_KEY=fw_KvkeXQmo8LctbP6xx8A5uA
-NVIDIA_API_KEY=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
-LLM_PROVIDER=nvidia
diff --git a/singular_ticker_causal/algorithms/CUTS_PLUS/cuts_plus.py b/singular_ticker_causal/algorithms/CUTS_PLUS/cuts_plus.py
deleted file mode 100644
index a99d2de63eb722af590896484a45edc9cb2e8318..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/algorithms/CUTS_PLUS/cuts_plus.py
+++ /dev/null
@@ -1,834 +0,0 @@
-import os
-from os.path import join as opj
-from os.path import dirname as opd
-
-import tqdm
-import numpy as np
-import argparse
-from omegaconf import OmegaConf
-from copy import deepcopy
-from einops import rearrange
-import torch
-from torch import nn
-import torch.nn.functional as F
-
-from .utils.gumbel_softmax import gumbel_softmax
-from .utils.misc import calc_and_log_metrics, log_time_series, plot_causal_matrix
-from .utils.opt_type import MultiCADopt
-from .utils.logger import MyLogger
-from .model.cuts_plus_net import CUTS_Plus_Net
-from causal_hierarchy.grouping import build_group_matrix
-from causal.cuts_plus.edge_controller import DualEdgeTemperatureController
-
-
-def accounting_prior_loss(
- G_sequence: torch.Tensor,
- G_p: torch.Tensor,
- lambda_s: float,
- lambda_d: float,
-) -> torch.Tensor:
- """Compute the combined sparseness + domain-fit prior loss for a lag-indexed graph sequence.
-
- Implements the plan §5.2 formula::
-
- p(G) ∝ exp( -λ_s ‖G_{1:L}‖_F² - λ_d ‖G_{1:L} - G^p_{1:L}‖_F² )
-
- Parameters
- ----------
- G_sequence : Tensor (L, D, D) or (D, D)
- Soft adjacency sequence. When 2-D it is treated as a single-lag graph
- and the sparseness / domain-fit are applied directly.
- G_p : Tensor (D, D)
- Static binary accounting prior mask. Broadcast across lags.
- lambda_s : float
- Sparseness regularisation coefficient.
- lambda_d : float
- Domain-fit (DuPont structural) regularisation coefficient.
-
- Returns
- -------
- Tensor scalar
- Combined prior loss term to be added to the CUTS+ objective.
- """
- if G_sequence.ndim == 2:
- G_sequence = G_sequence.unsqueeze(0) # treat as (1, D, D)
-
- # Broadcast static prior to (L, D, D)
- G_p_expanded = G_p.unsqueeze(0).expand_as(G_sequence)
-
- sparseness = lambda_s * torch.norm(G_sequence, p="fro") ** 2
- domain_fit = lambda_d * torch.norm(G_sequence - G_p_expanded, p="fro") ** 2
- return sparseness + domain_fit
-
-
-
-def plot_matrix(name, mat, log, log_step, vmin=None, vmax=None):
- if len(mat.shape) == 3:
- mat = np.max(mat, axis=-1)
- n, m = mat.shape
-
- # Show Discovered Graph (Probability)
- sub_cg = plot_causal_matrix(
- mat,
- figsize=[1.5*n, 1*n],
- show_text=False,
- vmin=vmin, vmax=vmax)
- log.log_figures(sub_cg, name=name, iters=log_step)
-
-
-def generate_indices(input_step, pred_step, t_length, block_size=None):
- if block_size is None:
- block_size = t_length
-
- offsets_in_block = np.arange(input_step, block_size-pred_step+1)
- assert t_length % block_size == 0, "t_length % block_size != 0"
- random_t_list = []
- for block_start in range(0, t_length, block_size):
- random_t_list += (offsets_in_block + block_start).tolist()
-
- np.random.shuffle(random_t_list)
- return random_t_list
-
-
-
-def batch_generater(data, observ_mask, bs, n_nodes, input_step, pred_step, block_size=None):
- t, n, d = data.shape
- first_sample_t = input_step
- random_t_list = generate_indices(input_step, pred_step, t_length=t, block_size=block_size)
-
- for batch_i in range(len(random_t_list) // bs):
- x = torch.zeros([bs, n_nodes, input_step, d]).to(data.device)
- y = torch.zeros([bs, n_nodes, pred_step, d]).to(data.device)
- t = torch.zeros([bs]).to(data.device).long()
- mask_x = torch.zeros([bs, n_nodes, input_step, d]).to(data.device)
- mask_y = torch.zeros([bs, n_nodes, pred_step, d]).to(data.device)
- for data_i in range(bs):
- data_t = random_t_list.pop()
- x[data_i, :, :, :] = rearrange(data[data_t-input_step : data_t, :], "t n d -> n t d")
- y[data_i, :, :, :] = rearrange(data[data_t : data_t+pred_step, :], "t n d -> n t d")
- t[data_i] = data_t
- mask_x[data_i, :, :, :] = rearrange(observ_mask[data_t-input_step : data_t, :], "t n d -> n t d")
- mask_y[data_i, :, :, :] = rearrange(observ_mask[data_t:data_t+pred_step, :], "t n d -> n t d")
-
- yield x, y, t, mask_x, mask_y
-
-
-
-
-
-class MultiCAD(object):
- def __init__(
- self,
- args: MultiCADopt.MultiCADargs,
- log,
- device="cuda",
- text_data=None,
- text_mask=None,
- G_prior=None,
- denoised_news=None,
- denoised_mask=None,
- ):
- self.log: MyLogger = log
- self.args = args
- self.device = device
-
- self.text_data = text_data.to(device) if text_data is not None else None
- self.text_mask = text_mask.to(device) if text_mask is not None else None
- self.denoised_news = denoised_news.to(device) if denoised_news is not None else None
- self.denoised_mask = denoised_mask.to(device) if denoised_mask is not None else None
- self.projector = None
- self.denoised_projector = None
-
- self.lambda_d = getattr(args, 'lambda_d', 1e-2)
- if G_prior is not None:
- self.G_prior = torch.from_numpy(G_prior).float().to(device)
- else:
- self.G_prior = torch.zeros(args.n_nodes, args.n_nodes).to(device)
- hard_edge_mask = getattr(args, "hard_edge_mask", None)
- if hard_edge_mask is None and bool(getattr(args, "use_hard_prior_edges", False)):
- hard_edge_mask = G_prior
- if hard_edge_mask is not None:
- self.hard_edge_mask_graph = torch.as_tensor(hard_edge_mask, dtype=torch.float32, device=device)
- else:
- self.hard_edge_mask_graph = torch.zeros(args.n_nodes, args.n_nodes, device=device)
- soft_edge_mask = getattr(args, "soft_edge_mask", None)
- if soft_edge_mask is not None:
- self.soft_edge_mask_graph = torch.as_tensor(soft_edge_mask, dtype=torch.float32, device=device)
- else:
- self.soft_edge_mask_graph = 1.0 - self.hard_edge_mask_graph
- self.hard_edge_tau = float(getattr(args, "hard_edge_tau", 0.02))
- self.hard_edge_trainable = bool(getattr(args, "hard_edge_trainable", False))
-
- # No embedding projector is required when text features are already
- # represented as low-dimensional sparse event tensors.
- self.projector = None
- self.denoised_projector = None
-
- self.fitting_model = CUTS_Plus_Net(self.args.n_nodes, in_ch=self.args.data_dim,
- n_layers=self.args.data_pred.gru_layers,
- hidden_ch=self.args.data_pred.mlp_hid,
- shared_weights_decoder=self.args.data_pred.shared_weights_decoder,
- concat_h=self.args.data_pred.concat_h,
- ).to(self.device)
-
- self.data_pred_loss = nn.MSELoss()
-
- params = list(self.fitting_model.parameters())
-
- self.data_pred_optimizer = torch.optim.Adam(
- params,
- lr=self.args.data_pred.lr_data_start,
- weight_decay=self.args.data_pred.weight_decay
- )
-
- if "every" in self.args.fill_policy:
- lr_schedule_length = int(self.args.fill_policy.split("_")[-1])
- else:
- lr_schedule_length = self.args.total_epoch
-
- gamma = (self.args.data_pred.lr_data_end / self.args.data_pred.lr_data_start) ** (1 / lr_schedule_length)
- self.data_pred_scheduler = torch.optim.lr_scheduler.StepLR(
- self.data_pred_optimizer, step_size=1, gamma=gamma)
-
- self.n_groups = self.args.n_groups
- print("n_groups: ", self.n_groups)
- if self.args.group_policy == "None":
- self.args.group_policy = None
- self.fixed_group_spec = self._resolve_fixed_group_spec()
- if self.fixed_group_spec is not None:
- self.n_groups = self.fixed_group_spec.n_groups
-
- end_tau, start_tau = self.args.graph_discov.end_tau, self.args.graph_discov.start_tau
- self.gumbel_tau_gamma = (end_tau / start_tau) ** (1 / self.args.total_epoch)
- self.gumbel_tau = start_tau
- self.start_tau = start_tau
- self.current_epoch = 0
- self.edge_controller = DualEdgeTemperatureController(
- G_prior=self.G_prior.detach().cpu().numpy(),
- tau_start=float(start_tau),
- tau_end=float(end_tau),
- tau_hard=self.hard_edge_tau,
- total_epochs=int(self.args.total_epoch),
- )
-
- end_lmd, start_lmd = self.args.graph_discov.lambda_s_end, self.args.graph_discov.lambda_s_start
- self.lambda_gamma = (end_lmd / start_lmd) ** (1 / self.args.total_epoch)
- self.lambda_s = start_lmd
-
- def set_graph_optimizer(self, epoch=None):
- if epoch == None:
- epoch = 0
-
- gamma = (self.args.graph_discov.lr_graph_end / self.args.graph_discov.lr_graph_start) ** (1 / self.args.total_epoch)
- self.graph_optimizer = torch.optim.Adam([self.GT], lr=self.args.graph_discov.lr_graph_start * gamma ** epoch)
- self.graph_scheduler = torch.optim.lr_scheduler.StepLR(self.graph_optimizer, step_size=1, gamma=gamma)
-
- def _resolve_fixed_group_spec(self):
- policy = getattr(self.args, "group_policy", None)
- if policy not in {"deterministic", "deterministic_sector", "deterministic_geography"}:
- return None
- assignments = getattr(self.args, "group_assignments", None)
- if assignments is None and policy == "deterministic_sector":
- ticker_list = getattr(self.args, "ticker_list", None)
- sector_map = getattr(self.args, "sector_map", None)
- if ticker_list is not None and sector_map is not None:
- assignments = [sector_map[ticker] for ticker in ticker_list]
- if assignments is None:
- raise ValueError("Deterministic grouping requires opt.group_assignments.")
- labels = getattr(self.args, "group_labels", None)
- return build_group_matrix(assignments, labels=labels)
-
- def _has_fixed_grouping(self) -> bool:
- return self.fixed_group_spec is not None
-
- def _init_random_gt(self, n_groups: int) -> torch.Tensor:
- return torch.ones((n_groups, self.args.n_nodes)) * 0.5 + torch.randn(n_groups, self.args.n_nodes) * 0.01
-
- def _build_prior_seed_logits(self, n_groups: int) -> torch.Tensor:
- gt_init = torch.full((n_groups, self.args.n_nodes), -2.0)
- if n_groups == self.args.n_nodes:
- grouped_prior = self.G_prior
- elif self._has_fixed_grouping():
- grouped_prior = torch.zeros(n_groups, self.args.n_nodes, device=self.device)
- for group_idx in range(n_groups):
- members = [
- idx for idx, assigned_group in enumerate(self.fixed_group_spec.assignments)
- if assigned_group == group_idx
- ]
- if members:
- grouped_prior[group_idx] = torch.max(self.G_prior[members], dim=0).values
- else:
- return self._init_random_gt(n_groups)
-
- gt_init[grouped_prior > 0.5] = 2.0
- gt_init += torch.randn_like(gt_init) * 0.05
- return gt_init
-
- def _build_graph_prob(self) -> torch.Tensor:
- return torch.einsum("nm,ml->nl", self.G, torch.sigmoid(self.GT))
-
- def _compose_graph_parts(self) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
- base_graph = self._build_graph_prob()
- soft_graph = base_graph * self.soft_edge_mask_graph
- if self.hard_edge_trainable:
- hard_graph = base_graph * self.hard_edge_mask_graph
- else:
- hard_graph = self.hard_edge_mask_graph
- effective_graph = torch.clamp(soft_graph + hard_graph, 0.0, 1.0)
- return soft_graph, hard_graph, effective_graph
-
- def _gumbel_sigmoid_sample(self, graph: torch.Tensor, batch_size: int, tau: float) -> torch.Tensor:
- prob = graph[None, :, :, None].expand(batch_size, -1, -1, -1)
- logits = torch.concat([prob, (1 - prob)], axis=-1)
- return gumbel_softmax(logits, tau=tau, hard=True)[:, :, :, 0]
-
- def _sample_graph_with_controller(self, graph: torch.Tensor) -> torch.Tensor:
- graph = torch.clamp(torch.nan_to_num(graph, nan=0.5), 1e-6, 1.0 - 1e-6)
- logits = torch.logit(graph)
- sampled = self.edge_controller.gumbel_sample(logits, epoch=self.current_epoch, hard=True)
- if not self.hard_edge_trainable and torch.any(self.hard_edge_mask_graph > 0):
- sampled = torch.clamp(sampled * self.soft_edge_mask_graph + self.hard_edge_mask_graph, 0.0, 1.0)
- return sampled[None].expand(self.args.batch_size, -1, -1)
-
- def _append_context(self, x, y, mask_x, mask_y, t, inp_step: int, pred_step: int):
- t_vals = t.cpu().tolist()
- if self.text_data is not None:
- tx = torch.stack([rearrange(self.text_data[ti - inp_step:ti], "t n d -> n t d") for ti in t_vals])
- ty = torch.stack([rearrange(self.text_data[ti:ti + pred_step], "t n d -> n t d") for ti in t_vals])
- d_text = self.text_data.shape[-1]
- if self.text_mask is not None:
- if self.text_mask.shape[-1] == 1:
- tmx = torch.stack([rearrange(self.text_mask[ti - inp_step:ti].expand(-1, -1, d_text), "t n d -> n t d") for ti in t_vals])
- tmy = torch.stack([rearrange(self.text_mask[ti:ti + pred_step].expand(-1, -1, d_text), "t n d -> n t d") for ti in t_vals])
- else:
- tmx = torch.stack([rearrange(self.text_mask[ti - inp_step:ti], "t n d -> n t d") for ti in t_vals])
- tmy = torch.stack([rearrange(self.text_mask[ti:ti + pred_step], "t n d -> n t d") for ti in t_vals])
- else:
- tmx = torch.ones_like(tx)
- tmy = torch.ones_like(ty)
- x = torch.cat([x, tx], dim=-1)
- y = torch.cat([y, ty], dim=-1)
- mask_x = torch.cat([mask_x, tmx], dim=-1)
- mask_y = torch.cat([mask_y, tmy], dim=-1)
-
- if self.denoised_news is not None:
- dx = torch.stack([rearrange(self.denoised_news[ti - inp_step:ti], "t n d -> n t d") for ti in t_vals])
- dy = torch.stack([rearrange(self.denoised_news[ti:ti + pred_step], "t n d -> n t d") for ti in t_vals])
- d_denoised = self.denoised_news.shape[-1]
- if self.denoised_mask is not None:
- dmask_src = self.denoised_mask
- if dmask_src.shape[-1] == 1:
- dmx = torch.stack([rearrange(dmask_src[ti - inp_step:ti].expand(-1, -1, d_denoised), "t n d -> n t d") for ti in t_vals])
- dmy = torch.stack([rearrange(dmask_src[ti:ti + pred_step].expand(-1, -1, d_denoised), "t n d -> n t d") for ti in t_vals])
- else:
- dmx = torch.stack([rearrange(dmask_src[ti - inp_step:ti], "t n d -> n t d") for ti in t_vals])
- dmy = torch.stack([rearrange(dmask_src[ti:ti + pred_step], "t n d -> n t d") for ti in t_vals])
- else:
- dmx = torch.ones_like(dx)
- dmy = torch.ones_like(dy)
- x = torch.cat([x, dx], dim=-1)
- y = torch.cat([y, dy], dim=-1)
- mask_x = torch.cat([mask_x, dmx], dim=-1)
- mask_y = torch.cat([mask_y, dmy], dim=-1)
-
- return x, y, mask_x, mask_y
-
- def _sample_graph_for_prediction(self, graph: torch.Tensor, batch_size: int) -> torch.Tensor:
- sample_matrix = graph[None].expand(batch_size, -1, -1)
- sample_matrix = torch.clamp(sample_matrix, 0.0, 1.0)
- sample_matrix = torch.nan_to_num(sample_matrix, nan=0.5)
- return torch.bernoulli(sample_matrix).float()
-
- def freeze_accounting_edges(self) -> None:
- """Stop gradient flow through GT logits that correspond to known prior edges.
-
- Called after epoch 10 (per plan §5.3). For positions where
- ``G_prior >= 0.5`` the logit is detached so that subsequent
- ``graph_optimizer.step()`` calls do not update those weights.
-
- The method operates in-place by replacing ``self.GT`` with a new
- ``nn.Parameter`` whose values at prior positions are detached
- constants while non-prior positions retain full gradient.
- """
- if not hasattr(self, "GT"):
- return
- with torch.no_grad():
- gt_data = self.GT.data.clone()
-
- # Build a mask aligned to GT shape (n_groups × n_nodes)
- n_groups, n_nodes = self.GT.shape
- if self.G_prior.shape == (n_nodes, n_nodes) and n_groups == n_nodes:
- # 1-to-1 mapping: prior mask applies directly
- prior_mask = (self.G_prior >= 0.5)
- elif self._has_fixed_grouping():
- # Map node-level prior to group-level: group is frozen if any member
- # has a prior edge from that group
- prior_mask = torch.zeros(n_groups, n_nodes, dtype=torch.bool, device=self.device)
- for group_idx in range(n_groups):
- members = [
- idx for idx, g in enumerate(self.fixed_group_spec.assignments)
- if g == group_idx
- ]
- if members:
- row_prior = self.G_prior[members].max(dim=0).values
- prior_mask[group_idx] = row_prior >= 0.5
- else:
- return # cannot determine mapping — skip freeze
-
- frozen_vals = gt_data[prior_mask].detach()
- new_gt = nn.Parameter(gt_data)
- # Freeze prior positions by zeroing their gradient contribution
- # via a register_hook that zeroes the grad at those positions.
- def _freeze_hook(grad: torch.Tensor) -> torch.Tensor:
- grad = grad.clone()
- grad[prior_mask] = 0.0
- return grad
-
- new_gt.register_hook(_freeze_hook)
- self.GT = new_gt
- self.set_graph_optimizer() # refresh optimizer to point at new GT
- n_frozen = int(prior_mask.sum().item())
- print(f"[freeze_accounting_edges] Froze {n_frozen} / {n_groups * n_nodes} GT logit positions.")
-
-
- def ticker_price_pred(self, x, y, mask_x, mask_y):
- bs, n, t, d = x.shape
- self.fitting_model.train()
- self.data_pred_optimizer.zero_grad()
-
- _, _, effective_graph = self._compose_graph_parts()
- graph_sampled = self._sample_graph_for_prediction(effective_graph, self.args.batch_size)
-
- y_pred = self.fitting_model(x, mask_x, graph_sampled)
-
- # print(y_pred.shape, y.shape, observ_mask.shape)
- loss = self.data_pred_loss(y * mask_y, y_pred * mask_y) / (torch.mean(mask_y) + 1e-8)
- loss.backward()
- self.data_pred_optimizer.step()
- return y_pred, loss
-
- def graph_discov(self, x, y, mask_x, mask_y):
- gn, n = self.GT.shape
- self.graph_optimizer.zero_grad()
- soft_graph, hard_graph, effective_graph = self._compose_graph_parts()
-
- graph_sampled = self._sample_graph_with_controller(effective_graph)
-
- loss_sparsity = torch.linalg.norm(soft_graph.flatten(), ord=1) / (n * n)
-
- y_pred = self.fitting_model(x, mask_x, graph_sampled)
-
- loss_data = self.data_pred_loss(y * mask_y, y_pred * mask_y) / (torch.mean(mask_y) + 1e-8)
-
- # DuPont structural prior penalty: push toward known edges, away from impossible ones
- loss_dupont = torch.linalg.norm((effective_graph - self.G_prior).flatten(), ord=2) ** 2 / (n * n)
-
- # L2 regularization on raw GT logits to prevent saturation to ±∞
- loss_l2_gt = torch.linalg.norm(self.GT.flatten(), ord=2) ** 2 / (gn * n)
- if torch.any(self.hard_edge_mask_graph > 0):
- hard_edge_density = effective_graph[self.hard_edge_mask_graph > 0].mean()
- else:
- hard_edge_density = torch.tensor(0.0, device=self.device)
-
- loss = (loss_sparsity * self.lambda_s
- + loss_data
- + self.lambda_d * loss_dupont
- + 1e-3 * loss_l2_gt) # small L2 keeps logits from drifting to ±∞
- loss.backward()
- self.graph_optimizer.step()
-
- return loss, loss_sparsity, loss_data, loss_dupont, hard_edge_density
-
-
-
- def train(self, data, observ_mask, original_data, true_cm=None):
-
- original_data = torch.from_numpy(original_data).float().to(self.device)
- observ_mask = torch.from_numpy(observ_mask).float().to(self.device)
- data = torch.from_numpy(data).float().to(self.device)
-
- if self.args.supervision_policy == "masked":
- print("Using masked supervision for data prediction...")
- elif self.args.supervision_policy == "full":
- print("Using full supervision for data prediction......")
- observ_mask = torch.ones_like(observ_mask)
- elif "masked_before" in self.args.supervision_policy:
- print(f"Using masked supervision for data prediction ({self.args.supervision_policy:s})......")
-
- price_pred_step = 0
- graph_discov_step = 0
- pbar = tqdm.tqdm(total=self.args.total_epoch)
- data_interp = deepcopy(data)
- original_mask = deepcopy(observ_mask)
- auc = 0
- _edges_frozen = False # track whether freeze_accounting_edges() has run
- for epoch_i in range(self.args.total_epoch):
- self.current_epoch = epoch_i
- # Phase 5 §5.3: freeze accounting edge logits after epoch 10
- if epoch_i == 10 and not _edges_frozen:
- self.freeze_accounting_edges()
- _edges_frozen = True
- if self._has_fixed_grouping():
- if epoch_i == 0:
- self.G = torch.from_numpy(self.fixed_group_spec.matrix).float().to(self.device)
- self.GT = nn.Parameter(self._build_prior_seed_logits(self.fixed_group_spec.n_groups).to(self.device))
- self.set_graph_optimizer(epoch_i)
- elif self.args.group_policy is not None:
- group_mul = int(self.args.group_policy.split("_")[1])
- group_every = int(self.args.group_policy.split("_")[3])
- if epoch_i % group_every == 0 and self.n_groups < self.args.n_nodes:
- if epoch_i != 0:
- self.n_groups *= group_mul
- if self.n_groups > self.args.n_nodes:
- self.n_groups = self.args.n_nodes
-
- self.G = torch.zeros([self.args.n_nodes, self.n_groups]).to(self.device)
-
- for i in range(0, self.n_groups):
- for j in range(0, self.args.n_nodes // self.n_groups):
- self.G[i*(self.args.n_nodes // self.n_groups) + j, i] = 1
- for k in range(i*(self.args.n_nodes // self.n_groups) + j, self.args.n_nodes):
- self.G[k, i] = 1
-
- if hasattr(self, "GT"):
- GT_init = torch.sigmoid(self.GT).detach().cpu().repeat_interleave(group_mul, 0)[:self.n_groups, :]
- GT_init = 1 - (1 - GT_init)**(1 / group_mul)
- else:
- GT_init = self._init_random_gt(self.n_groups)
-
- self.GT = nn.Parameter(GT_init.to(self.device))
-
- self.set_graph_optimizer(epoch_i)
- elif epoch_i == 0 and self.n_groups == self.args.n_nodes:
- self.G = torch.eye(self.args.n_nodes).to(self.device)
- # Add small noise to break symmetry — identical init → identical gradients
- GT_init = torch.ones((self.n_groups, self.args.n_nodes))*0.5 + torch.randn(self.n_groups, self.args.n_nodes)*0.01
- self.GT = nn.Parameter(GT_init.to(self.device))
- self.set_graph_optimizer(epoch_i)
- else:
- if epoch_i == 0:
- self.n_groups = self.args.n_nodes
- self.G = torch.eye(self.args.n_nodes).to(self.device)
- GT_init = self._build_prior_seed_logits(self.n_groups)
- self.GT = nn.Parameter(GT_init.to(self.device))
- self.set_graph_optimizer(epoch_i)
-
-
- if "every" in self.args.fill_policy:
- update_every = int(self.args.fill_policy.split("_")[-1])
- if (epoch_i+1) % update_every == 0:
- data = data_pred
- print("Update data!")
- # self.graph_optimizer.param_groups[0]['lr'] = self.args.graph_discov.lr_graph_start
- self.data_pred_optimizer.param_groups[0]['lr'] = self.args.data_pred.lr_data_start
- observ_mask = torch.ones_like(original_mask)
- elif "rate" in self.args.fill_policy:
- update_rate = float(self.args.fill_policy.split("_")[1])
- update_after = int(self.args.fill_policy.split("_")[3])
- if epoch_i+1 > update_after:
- if epoch_i == update_after:
- print("Data update started!")
- data = data * (1 - update_rate) + data_pred * update_rate
- else:
- # no data update
- pass
-
- if "masked_before" in self.args.supervision_policy:
- masked_before = int(self.args.supervision_policy.split("_")[2])
- if epoch_i == masked_before:
- print("Using full supervision for data prediction......")
- observ_mask = torch.ones_like(original_mask)
- self.gumbel_tau = self.start_tau
-
- # Data Prediction
- if hasattr(self.args, "data_pred"):
- if hasattr(self.args, "block_size"):
- block_size = self.args.block_size
- else:
- block_size = None
-
- # Always use tech-only data for the batch generator.
- # If a projector exists, text is projected fresh INSIDE each batch
- # to avoid stale computation graphs after optimizer.step().
- batch_gen = batch_generater(data, observ_mask,
- bs=self.args.batch_size,
- n_nodes=self.args.n_nodes,
- input_step=self.args.input_step,
- pred_step=self.args.data_pred.pred_step,
- block_size=block_size)
- batch_gen = list(batch_gen)
-
- data_pred = data.clone().detach() # tech-only predictions
- data_pred_all = data.clone().detach()
- d_tech = data.shape[-1]
- inp_step = self.args.input_step
- pred_step = self.args.data_pred.pred_step
-
- for x, y, t, mask_x, mask_y in batch_gen:
- price_pred_step += self.args.batch_size
-
- x, y, mask_x, mask_y = self._append_context(x, y, mask_x, mask_y, t, inp_step, pred_step)
-
- y_pred, loss = self.ticker_price_pred(x, y, mask_x, mask_y)
- # Map back only the tech portion
- data_pred[t] = (y_pred*(1-mask_y) + y*mask_y).clone().detach()[:,:,0,:d_tech]
- data_pred_all[t] = y_pred.clone().detach()[:,:,0,:d_tech]
- self.log.log_metrics({"ticker_price_pred/pred_loss": loss.item()}, price_pred_step)
- pbar.set_postfix_str(f"S1 loss={loss.item():.2f}, spr=IDLE, auc={auc:.4f}")
-
- current_data_pred_lr = self.data_pred_optimizer.param_groups[0]['lr']
- self.log.log_metrics({"graph_discov/lr": current_data_pred_lr}, price_pred_step)
- self.data_pred_scheduler.step()
- mse_pred_to_original = self.data_pred_loss(original_data, data_pred)
- mse_interp_to_original = self.data_pred_loss(original_data, data_interp)
-
- self.log.log_metrics({"ticker_price_pred/mse_pred_to_original": mse_pred_to_original,
- "ticker_price_pred/mse_interp_to_original": mse_interp_to_original}, price_pred_step)
-
- # Graph Discovery
- if hasattr(self.args, "graph_discov"):
- for x, y, t, mask_x, mask_y in batch_gen:
- graph_discov_step += self.args.batch_size
- if hasattr(self.args, "disable_graph") and self.args.disable_graph:
- pass
- else:
- x, y, mask_x, mask_y = self._append_context(x, y, mask_x, mask_y, t, inp_step, pred_step)
-
- loss, loss_sparsity, loss_data, loss_dupont, hard_edge_density = self.graph_discov(x, y, mask_x, mask_y)
- self.log.log_metrics({"graph_discov/sparsity_loss": loss_sparsity.item(),
- "graph_discov/data_loss": loss_data.item(),
- "graph_discov/prior_loss": loss_dupont.item(),
- "graph_discov/hard_edge_density": hard_edge_density.item(),
- "graph_discov/total_loss": loss.item()}, graph_discov_step)
- pbar.set_postfix_str(f"S2 loss={loss_data.item():.2f}, spr={loss_sparsity.item():.2f}, auc={auc:.4f}")
-
- self.graph_scheduler.step()
- # self.group_scheduler.step()
- current_graph_disconv_lr = self.graph_optimizer.param_groups[0]['lr']
- self.log.log_metrics({"graph_discov/lr": current_graph_disconv_lr}, graph_discov_step)
- self.log.log_metrics({"graph_discov/tau": self.gumbel_tau}, graph_discov_step)
- self.gumbel_tau *= self.gumbel_tau_gamma
- self.lambda_s *= self.lambda_gamma
-
- pbar.update(1)
-
- plot_roc = False
-
- G_prob = self.G.detach().cpu().numpy()
- GT_prob = self.GT.detach().cpu().numpy()
- # Apply sigmoid to match training forward pass (ticker_price_pred/graph_discov use sigmoid)
- GT_prob_sigmoid = 1 / (1 + np.exp(-GT_prob))
- Graph = np.einsum("nm,ml->nl", G_prob, GT_prob_sigmoid)
- if np.any(self.hard_edge_mask_graph.detach().cpu().numpy() > 0):
- if self.hard_edge_trainable:
- Graph = Graph * self.soft_edge_mask_graph.detach().cpu().numpy() + Graph * self.hard_edge_mask_graph.detach().cpu().numpy()
- else:
- Graph = Graph * self.soft_edge_mask_graph.detach().cpu().numpy() + self.hard_edge_mask_graph.detach().cpu().numpy()
- Graph = np.clip(Graph, 0.0, 1.0)
-
-
- if (epoch_i+1) % self.args.show_graph_every == 0:
- avg_mask = np.mean(observ_mask.cpu().numpy(), axis=(0,2))
- if np.min(avg_mask) < 1:
- time_series_idx = int(np.argwhere(avg_mask < 1)[0, 0])
- else:
- time_series_idx = 0
- d_tech = original_data.shape[-1]
- log_time_series(
- original_data.cpu()[-100:,time_series_idx],
- data_interp.cpu()[-100:,time_series_idx],
- data_pred_all.cpu()[-100:,time_series_idx, :d_tech],
- log=self.log, log_step=price_pred_step
- )
-
- plot_matrix("G", G_prob, self.log, graph_discov_step, vmin=0, vmax=1)
- plot_matrix("GT", GT_prob, self.log, graph_discov_step, vmin=0, vmax=1)
- plot_matrix("Graph", Graph, self.log, graph_discov_step, vmin=0, vmax=1)
- np.save(os.path.join(self.log.log_dir, 'Graph.npy'), Graph)
- plot_roc = True
-
- # Show TPR FPR AUC ROC
- if true_cm is not None:
- Graph = rearrange(Graph, "n m -> m n")
- auc = calc_and_log_metrics(Graph, true_cm, self.log, graph_discov_step, plot_roc=plot_roc)
-
- return Graph
-
-
-def prepross_data(data):
- T, N, D = data.shape
- new_data = np.zeros_like(data, dtype=float)
- for i in range(N):
- node = data[:,i,:]
- std = np.std(node)
- # Guard against zero-std (constant) columns to prevent NaN from 0/0
- new_data[:,i,:] = (node - np.mean(node)) / (std + 1e-8)
- # Replace any residual NaN/Inf (e.g. from upstream data issues) with 0
- new_data = np.nan_to_num(new_data, nan=0.0, posinf=0.0, neginf=0.0)
- return new_data
-
-
-def main(
- data,
- mask,
- true_cm,
- opt,
- log,
- device="cuda",
- text_data=None,
- text_mask=None,
- G_prior=None,
- denoised_news=None,
- denoised_mask=None,
-):
- if opt.n_nodes == "auto":
- opt.n_nodes = data.shape[1]
-
- if len(data.shape) == 2:
- data = data[:,:,None]
- mask = mask[:,:,None]
- data = prepross_data(data)
-
- if text_data is not None:
- text_data_torch = torch.from_numpy(text_data).float()
- text_mask_torch = torch.from_numpy(text_mask).float()
- else:
- text_data_torch = None
- text_mask_torch = None
- if denoised_news is not None:
- denoised_news_torch = torch.from_numpy(denoised_news).float()
- denoised_mask_torch = torch.from_numpy(denoised_mask).float() if denoised_mask is not None else None
- else:
- denoised_news_torch = None
- denoised_mask_torch = None
-
- effective_dim = opt.data_dim
- projector_output_dim = getattr(opt, 'projector_output_dim', 16)
- if text_data is not None:
- effective_dim += projector_output_dim
- if denoised_news is not None:
- effective_dim += projector_output_dim
- opt.data_dim = effective_dim
-
- multicad = MultiCAD(
- opt,
- log,
- device=device,
- text_data=text_data_torch,
- text_mask=text_mask_torch,
- G_prior=G_prior,
- denoised_news=denoised_news_torch,
- denoised_mask=denoised_mask_torch,
- )
- max_refuter_retries = 3
- refuter_retries = 0
- falsified = True
-
- while falsified and refuter_retries <= max_refuter_retries:
- Graph = multicad.train(data, mask, data, true_cm)
-
- # Run Refuter Validation
- try:
- from singular_ticker_causal.causal_inference.causal_model import StructuralCausalModel
- from singular_ticker_causal.causal_inference.query_engine import CausalQueryEngine
-
- data_cpu = data.cpu().numpy() if isinstance(data, torch.Tensor) else data
- n_nodes = Graph.shape[0]
- nodes = [f"N_{i}" for i in range(n_nodes)]
- adj_mask = G_prior if G_prior is not None else np.zeros_like(Graph)
-
- scm = StructuralCausalModel(
- nodes=nodes,
- adj=Graph,
- adjacency_mask=adj_mask,
- data_tech=data_cpu,
- lag=1
- ).fit()
-
- pywhyllm_enabled = os.environ.get("PYWHYLLM_ENABLED", "").lower() in {"1", "true", "yes", "on"}
- pywhyllm_max_edges = int(os.environ.get("PYWHYLLM_REFUTER_MAX_EDGES", "3"))
- engine = CausalQueryEngine(
- scm,
- data_tech=data_cpu,
- pywhyllm_enabled=pywhyllm_enabled,
- )
-
- falsified = False
- failed_edges = []
-
- edges_to_check = []
- for i in range(n_nodes):
- for j in range(n_nodes):
- if scm.dag_adj[i, j] and (adj_mask[i, j] == 0):
- edges_to_check.append((nodes[i], nodes[j], float(Graph[i, j])))
- edges_to_check.sort(key=lambda edge: abs(edge[2]), reverse=True)
-
- validation_reports = []
- for treatment, outcome, _score in edges_to_check[:pywhyllm_max_edges]:
- if pywhyllm_enabled:
- res = engine.validate_with_pywhyllm_and_dowhy(
- treatment,
- outcome,
- max_edges=pywhyllm_max_edges,
- )
- else:
- res = engine.validate_with_dowhy(treatment, outcome)
- validation_reports.append({"edge": [treatment, outcome], "validation": res})
- if res.get("falsified"):
- falsified = True
- failed_edges.append((treatment, outcome, res))
- break
-
- if falsified:
- print(f"[Refuter] Graph failed refutation on edges: {failed_edges}. Applying penalty and retrying...")
- multicad.args.total_epoch = 100
- multicad.lambda_s *= 1.5
-
- if hasattr(multicad, "GT"):
- multicad.GT.data = multicad._build_prior_seed_logits(multicad.n_groups).to(device)
- refuter_retries += 1
-
- import json
- artifact = {
- "failed_edges": [
- {"treatment": edge[0], "outcome": edge[1], "validation": edge[2]}
- for edge in failed_edges
- ],
- "validation_reports": validation_reports,
- "lambda_s_new": float(multicad.lambda_s),
- "retry_epoch": 100,
- "adj": Graph.tolist()
- }
- with open(os.path.join(log.log_dir, f"refuter_failed_artifact_retry_{refuter_retries}.json"), "w") as f:
- json.dump(artifact, f)
- else:
- print("[Refuter] Graph passed refutation or no testable edges.")
-
- except Exception as e:
- print(f"[Refuter] Validation error: {e}. Bypassing refuter.")
- falsified = False
-
- return Graph
-
-
-if __name__ == "__main__":
- os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID"
-
- parser = argparse.ArgumentParser(description='Batch Compress')
- parser.add_argument('-opt', type=str, default=opj(opd(__file__),
- 'opt/multi_cad_lorenz.yaml'), help='yaml file path')
- parser.add_argument('-g', help='availabel gpu list', default='2', type=str)
- parser.add_argument('-debug', action='store_true')
- parser.add_argument('-log', action='store_true')
- args = parser.parse_args()
-
- if args.g == "mps":
- os.environ["PYTORCH_ENABLE_MPS_FALLBACK"] = "1"
- device = "mps"
- elif args.g == "cpu":
- device = "cpu"
- else:
- os.environ["CUDA_VISIBLE_DEVICES"] = args.g
- device = "cuda"
-
- main(OmegaConf.load(args.opt), device=device)
diff --git a/singular_ticker_causal/algorithms/CUTS_PLUS/data/__init__.py b/singular_ticker_causal/algorithms/CUTS_PLUS/data/__init__.py
deleted file mode 100644
index e69de29bb2d1d6434b8b29ae775ad8c2e48c5391..0000000000000000000000000000000000000000
diff --git a/singular_ticker_causal/algorithms/CUTS_PLUS/data/generate_data_mod.py b/singular_ticker_causal/algorithms/CUTS_PLUS/data/generate_data_mod.py
deleted file mode 100644
index 131c013b4e3ad494e42fd47607323cfc29310d6e..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/algorithms/CUTS_PLUS/data/generate_data_mod.py
+++ /dev/null
@@ -1,430 +0,0 @@
-import numpy as np
-from collections import defaultdict
-
-def check_stationarity(links):
- """Returns stationarity according to a unit root test
-
- Assuming a Gaussian Vector autoregressive process
-
- Three conditions are necessary for stationarity of the VAR(p) model:
- - Absence of mean shifts;
- - The noise vectors are identically distributed;
- - Stability condition on Phi(t-1) coupling matrix (stabmat) of VAR(1)-version of VAR(p).
- """
-
-
- N = len(links)
- # Check parameters
- max_lag = 0
-
- for j in range(N):
- for link_props in links[j]:
- var, lag = link_props[0]
- # coeff = link_props[1]
- # coupling = link_props[2]
-
- max_lag = max(max_lag, abs(lag))
-
- graph = np.zeros((N,N,max_lag))
- couplings = []
-
- for j in range(N):
- for link_props in links[j]:
- var, lag = link_props[0]
- coeff = link_props[1]
- coupling = link_props[2]
- if abs(lag) > 0:
- graph[j,var,abs(lag)-1] = coeff
- couplings.append(coupling)
-
- stabmat = np.zeros((N*max_lag,N*max_lag))
- index = 0
-
- for i in range(0,N*max_lag,N):
- stabmat[:N,i:i+N] = graph[:,:,index]
- if index < max_lag-1:
- stabmat[i+N:i+2*N,i:i+N] = np.identity(N)
- index += 1
-
- eig = np.linalg.eig(stabmat)[0]
- # print "----> maxeig = ", np.abs(eig).max()
- if np.all(np.abs(eig) < 1.):
- stationary = True
- else:
- stationary = False
-
- if len(eig) == 0:
- return stationary, 0.
- else:
- return stationary, np.abs(eig).max()
-
-
-class Graph():
- def __init__(self,vertices):
- self.graph = defaultdict(list)
- self.V = vertices
-
- def addEdge(self,u,v):
- self.graph[u].append(v)
-
- def isCyclicUtil(self, v, visited, recStack):
-
- # Mark current node as visited and
- # adds to recursion stack
- visited[v] = True
- recStack[v] = True
-
- # Recur for all neighbours
- # if any neighbour is visited and in
- # recStack then graph is cyclic
- for neighbour in self.graph[v]:
- if visited[neighbour] == False:
- if self.isCyclicUtil(neighbour, visited, recStack) == True:
- return True
- elif recStack[neighbour] == True:
- return True
-
- # The node needs to be poped from
- # recursion stack before function ends
- recStack[v] = False
- return False
-
- # Returns true if graph is cyclic else false
- def isCyclic(self):
- visited = [False] * self.V
- recStack = [False] * self.V
- for node in range(self.V):
- if visited[node] == False:
- if self.isCyclicUtil(node,visited,recStack) == True:
- return True
- return False
-
- # A recursive function used by topologicalSort
- def topologicalSortUtil(self,v,visited,stack):
-
- # Mark the current node as visited.
- visited[v] = True
-
- # Recur for all the vertices adjacent to this vertex
- for i in self.graph[v]:
- if visited[i] == False:
- self.topologicalSortUtil(i,visited,stack)
-
- # Push current vertex to stack which stores result
- stack.insert(0,v)
-
- # The function to do Topological Sort. It uses recursive
- # topologicalSortUtil()
- def topologicalSort(self):
- # Mark all the vertices as not visited
- visited = [False]*self.V
- stack =[]
-
- # Call the recursive helper function to store Topological
- # Sort starting from all vertices one by one
- for i in range(self.V):
- if visited[i] == False:
- self.topologicalSortUtil(i,visited,stack)
-
- return stack
-
-def generate_nonlinear_contemp_timeseries(links, T, noises=None, random_state=None):
-
- if random_state is None:
- random_state = np.random
-
- # links must be {j:[((i, -tau), func), ...], ...}
- # coeff is coefficient
- # func is a function f(x) that becomes linear ~x in limit
- # noises is a random_state.___ function
- N = len(links.keys())
- if noises is None:
- noises = [random_state.randn for j in range(N)]
-
- if N != max(links.keys())+1 or N != len(noises):
- raise ValueError("links and noises keys must match N.")
-
- # Check parameters
- max_lag = 0
- contemp = False
- contemp_dag = Graph(N)
- causal_order = list(range(N))
- for j in range(N):
- for link_props in links[j]:
- var, lag = link_props[0]
- coeff = link_props[1]
- func = link_props[2]
- if lag == 0: contemp = True
- if var not in range(N):
- raise ValueError("var must be in 0..{}.".format(N-1))
- if 'float' not in str(type(coeff)):
- raise ValueError("coeff must be float.")
- if lag > 0 or type(lag) != int:
- raise ValueError("lag must be non-positive int.")
- max_lag = max(max_lag, abs(lag))
-
- # Create contemp DAG
- if var != j and lag == 0:
- contemp_dag.addEdge(var, j)
- # a, b = causal_order.index(var), causal_order.index(j)
- # causal_order[b], causal_order[a] = causal_order[a], causal_order[b]
-
- if contemp_dag.isCyclic() == 1:
- raise ValueError("Contemporaneous links must not contain cycle.")
-
- causal_order = contemp_dag.topologicalSort()
-
- transient = int(.2*T)
-
- X = np.zeros((T+transient, N), dtype='float32')
- for j in range(N):
- X[:, j] = noises[j](T+transient)
-
- for t in range(max_lag, T+transient):
- for j in causal_order:
- for link_props in links[j]:
- var, lag = link_props[0]
- # if abs(lag) > 0:
- coeff = link_props[1]
- func = link_props[2]
-
- X[t, j] += coeff * func(X[t + lag, var])
-
- X = X[transient:]
-
- if (check_stationarity(links)[0] == False or
- np.any(np.isnan(X)) or
- np.any(np.isinf(X)) or
- # np.max(np.abs(X)) > 1.e4 or
- np.any(np.abs(np.triu(np.corrcoef(X, rowvar=0), 1)) > 0.999)):
- nonstationary = True
- else:
- nonstationary = False
-
- return X, nonstationary
-
-
-def generate_random_contemp_model(N, L,
- coupling_coeffs,
- coupling_funcs,
- auto_coeffs,
- tau_max,
- contemp_fraction=0.,
- # num_trials=1000,
- random_state=None):
-
- def lin(x): return x
-
- if random_state is None:
- random_state = np.random
-
- # print links
- a_len = len(auto_coeffs)
- if type(coupling_coeffs) == float:
- coupling_coeffs = [coupling_coeffs]
- c_len = len(coupling_coeffs)
- func_len = len(coupling_funcs)
-
- if tau_max == 0:
- contemp_fraction = 1.
-
- if contemp_fraction > 0.:
- contemp = True
- L_lagged = int((1.-contemp_fraction)*L)
- L_contemp = L - L_lagged
- if L==1:
- # Randomly assign a lagged or contemp link
- L_lagged = random_state.randint(0,2)
- L_contemp = int(L_lagged == False)
-
- else:
- contemp = False
- L_lagged = L
- L_contemp = 0
-
-
- # for ir in range(num_trials):
-
- # Random order
- causal_order = list(random_state.permutation(N))
-
- links = dict([(i, []) for i in range(N)])
-
- # Generate auto-dependencies at lag 1
- if tau_max > 0:
- for i in causal_order:
- a = auto_coeffs[random_state.randint(0, a_len)]
-
- if a != 0.:
- links[i].append(((int(i), -1), float(a), lin))
-
- chosen_links = []
- # Create contemporaneous DAG
- contemp_links = []
- for l in range(L_contemp):
-
- cause = random_state.choice(causal_order[:-1])
- effect = random_state.choice(causal_order)
- while (causal_order.index(cause) >= causal_order.index(effect)
- or (cause, effect) in chosen_links):
- cause = random_state.choice(causal_order[:-1])
- effect = random_state.choice(causal_order)
-
- contemp_links.append((cause, effect))
- chosen_links.append((cause, effect))
-
- # Create lagged links (can be cyclic)
- lagged_links = []
- for l in range(L_lagged):
-
- cause = random_state.choice(causal_order)
- effect = random_state.choice(causal_order)
- while (cause, effect) in chosen_links or cause == effect:
- cause = random_state.choice(causal_order)
- effect = random_state.choice(causal_order)
-
- lagged_links.append((cause, effect))
- chosen_links.append((cause, effect))
-
- # print(chosen_links)
- # print(contemp_links)
- for (i, j) in chosen_links:
-
- # Choose lag
- if (i, j) in contemp_links:
- tau = 0
- else:
- tau = int(random_state.randint(1, tau_max+1))
- # print tau
- # CHoose coupling
- c = float(coupling_coeffs[random_state.randint(0, c_len)])
- if c != 0:
- func = coupling_funcs[random_state.randint(0, func_len)]
-
- links[j].append(((int(i), -tau), c, func))
-
- # # Stationarity check assuming model with linear dependencies at least for large x
- # # if check_stationarity(links)[0]:
- # # return links
- # X, nonstat = generate_nonlinear_contemp_timeseries(links,
- # T=10000, noises=None, random_state=None)
- # if nonstat == False:
- # return links
- # else:
- # print("Trial %d: Not a stationary model" % ir)
-
-
- # print("No stationary models found in {} trials".format(num_trials))
- return links
-
-def generate_logistic_maps(N, T, links, noise_lev):
-
- # Check parameters
- # contemp = False
- max_lag = 0
- for j in range(N):
- for link_props in links[j]:
- var, lag = link_props[0]
- max_lag = max(max_lag, abs(lag))
-
- transient = int(.2*T)
-
- # Chaotic logistic map parameter
- r = 4.
-
- X = np.random.rand(T+transient, N)
-
- for t in range(max_lag, T+transient):
- for j in range(N):
- added_input = 0.
- for link_props in links[j]:
- var, lag = link_props[0]
- if var != j and abs(lag) > 0:
- coeff = link_props[1]
- coupling = link_props[2]
- added_input += coeff*X[t - abs(lag), var]
-
- X[t, j] = (X[t-1, j] * (r - r*X[t-1, j] - added_input + noise_lev*np.random.rand())) % 1
- #func(coeff, X[t+lag, var], coupling)
-
- X = X[transient:]
-
- if np.any(np.abs(X) == np.inf) or np.any(X == np.nan):
- raise ValueError("Data divergent")
- return X
-
-
-
-def weighted_avg_and_std(values, axis, weights):
- """Returns the weighted average and standard deviation.
-
- Parameters
- ---------
- values : array
- Data array of shape (time, variables).
-
- axis : int
- Axis to average/std about
-
- weights : array
- Weight array of shape (time, variables).
-
- Returns
- -------
- (average, std) : tuple of arrays
- Tuple of weighted average and standard deviation along axis.
- """
-
- values[np.isnan(values)] = 0.
- average = np.ma.average(values, axis=axis, weights=weights)
- variance = np.sum(weights * (values - np.expand_dims(average, axis)
- ) ** 2, axis=axis) / weights.sum(axis=axis)
-
- return (average, np.sqrt(variance))
-
-def time_bin_with_mask(data, time_bin_length, sample_selector=None):
- """Returns time binned data where only about non-masked values is averaged.
-
- Parameters
- ----------
- data : array
- Data array of shape (time, variables).
-
- time_bin_length : int
- Length of time bin.
-
- mask : bool array, optional (default: None)
- Data mask where True labels masked samples.
-
- Returns
- -------
- (bindata, T) : tuple of array and int
- Tuple of time-binned data array and new length of array.
- """
-
- T = len(data)
-
- time_bin_length = int(time_bin_length)
-
- if sample_selector is None:
- sample_selector = np.ones(data.shape)
-
- if np.ndim(data) == 1.:
- data.shape = (T, 1)
- sample_selector.shape = (T, 1)
-
- bindata = np.zeros(
- (T // time_bin_length,) + data.shape[1:], dtype="float32")
- for index, i in enumerate(range(0, T - time_bin_length + 1,
- time_bin_length)):
- # print weighted_avg_and_std(fulldata[i:i+time_bin_length], axis=0,
- # weights=sample_selector[i:i+time_bin_length])[0]
- bindata[index] = weighted_avg_and_std(data[i:i + time_bin_length],
- axis=0,
- weights=sample_selector[i:i +
- time_bin_length])[0]
-
- T, grid_size = bindata.shape
-
- return (bindata.squeeze(), T)
-
diff --git a/singular_ticker_causal/algorithms/CUTS_PLUS/data/simu_data.py b/singular_ticker_causal/algorithms/CUTS_PLUS/data/simu_data.py
deleted file mode 100644
index 8f9c88dd36ee755a6429bf66f3b89a78ddacb716..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/algorithms/CUTS_PLUS/data/simu_data.py
+++ /dev/null
@@ -1,313 +0,0 @@
-import os
-import sys
-from os.path import join as opj
-sys.path.append(opj(os.getcwd(), "../"))
-sys.path.append(os.getcwd())
-
-import csv
-import torch
-import scipy
-from .generate_data_mod import generate_random_contemp_model, generate_nonlinear_contemp_timeseries
-import numpy as np
-from scipy.integrate import odeint
-
-
-######################################
-# Function for loading input data
-######################################
-def loadTrainingData(inputDataFilePath, device):
-
- # Load and parse input data (create batch data)
- inpData = torch.load(inputDataFilePath)
- Xtrain = torch.zeros(inpData['TsData'].shape[1], inpData['TsData'].shape[0], requires_grad = False, device=device)
- Xtrain1 = inpData['TsData'].t()
- Xtrain.data[:,:] = Xtrain1.data[:,:]
-
- return Xtrain
-
-#######################################################
-# Function for reading ground truth network from file
-#######################################################
-def loadTrueNetwork(inputFilePath, networkSize):
-
- with open(inputFilePath) as tsvin:
- reader = csv.reader(tsvin, delimiter='\t')
- numrows = 0
- for row in reader:
- numrows = numrows + 1
-
- network = np.zeros((numrows,2),dtype=np.int16)
- with open(inputFilePath) as tsvin:
- reader = csv.reader(tsvin, delimiter='\t')
- rowcounter = 0
- for row in reader:
- network[rowcounter][0] = int(row[0][1:])
- network[rowcounter][1] = int(row[1][1:])
- rowcounter = rowcounter + 1
-
- Gtrue = np.zeros((networkSize,networkSize), dtype=np.int16)
- for row in range(0,len(network),1):
- Gtrue[network[row][1]-1][network[row][0]-1] = 1
-
- return Gtrue
-
-
-def load_dream_data(dataset_id):
- device = "cpu"
-
- if(dataset_id == 0):
- InputDataFilePath = "causal_discov_sota/SRU_for_GCI/data/dream3/Dream3TensorData/Size100Ecoli1.pt"
- RefNetworkFilePath = "causal_discov_sota/SRU_for_GCI/data/dream3/TrueGeneNetworks/InSilicoSize100-Ecoli1.tsv"
- elif(dataset_id == 1):
- InputDataFilePath = "causal_discov_sota/SRU_for_GCI/data/dream3/Dream3TensorData/Size100Ecoli2.pt"
- RefNetworkFilePath = "causal_discov_sota/SRU_for_GCI/data/dream3/TrueGeneNetworks/InSilicoSize100-Ecoli2.tsv"
- elif(dataset_id == 2):
- InputDataFilePath = "causal_discov_sota/SRU_for_GCI/data/dream3/Dream3TensorData/Size100Yeast1.pt"
- RefNetworkFilePath = "causal_discov_sota/SRU_for_GCI/data/dream3/TrueGeneNetworks/InSilicoSize100-Yeast1.tsv"
- elif(dataset_id == 3):
- InputDataFilePath = "causal_discov_sota/SRU_for_GCI/data/dream3/Dream3TensorData/Size100Yeast2.pt"
- RefNetworkFilePath = "causal_discov_sota/SRU_for_GCI/data/dream3/TrueGeneNetworks/InSilicoSize100-Yeast2.tsv"
- elif(dataset_id == 4):
- InputDataFilePath = "causal_discov_sota/SRU_for_GCI/data/dream3/Dream3TensorData/Size100Yeast3.pt"
- RefNetworkFilePath = "causal_discov_sota/SRU_for_GCI/data/dream3/TrueGeneNetworks/InSilicoSize100-Yeast3.tsv"
- elif(dataset_id == 5):
- InputDataFilePath = "causal_discov_sota/SRU_for_GCI/data/dream3/Dream3TensorData/Size10Ecoli1.pt"
- RefNetworkFilePath = "causal_discov_sota/SRU_for_GCI/data/dream3/TrueGeneNetworks/InSilicoSize10-Ecoli1.tsv"
- elif(dataset_id == 6):
- InputDataFilePath = "causal_discov_sota/SRU_for_GCI/data/dream3/Dream3TensorData/Size10Ecoli2.pt"
- RefNetworkFilePath = "causal_discov_sota/SRU_for_GCI/data/dream3/TrueGeneNetworks/InSilicoSize10-Ecoli2.tsv"
- elif(dataset_id == 7):
- InputDataFilePath = "causal_discov_sota/SRU_for_GCI/data/dream3/Dream3TensorData/Size10Yeast1.pt"
- RefNetworkFilePath = "causal_discov_sota/SRU_for_GCI/data/dream3/TrueGeneNetworks/InSilicoSize10-Yeast1.tsv"
- elif(dataset_id == 8):
- InputDataFilePath = "causal_discov_sota/SRU_for_GCI/data/dream3/Dream3TensorData/Size10Yeast2.pt"
- RefNetworkFilePath = "causal_discov_sota/SRU_for_GCI/data/dream3/TrueGeneNetworks/InSilicoSize10-Yeast2.tsv"
- elif(dataset_id == 9):
- InputDataFilePath = "causal_discov_sota/SRU_for_GCI/data/dream3/Dream3TensorData/Size10Yeast3.pt"
- RefNetworkFilePath = "causal_discov_sota/SRU_for_GCI/data/dream3/TrueGeneNetworks/InSilicoSize10-Yeast3.tsv"
- else:
- print("Error while loading gene training data")
-
- Xtrain = loadTrainingData(InputDataFilePath, device)
- n = Xtrain.shape[0]
- Gref = loadTrueNetwork(RefNetworkFilePath, n)
-
- Xtrain = Xtrain.numpy().T
- # Gref = Gref.T
-
- return Xtrain, Gref
-
-
-
-def links_to_matrix(links):
- N = len(links)
- cm = np.zeros([N, N])
- for i, effect_node in links.items():
- for (j, _), _, _ in effect_node:
- cm[i, j] += 1
- return cm
-
-
-class noise_model:
- def __init__(self, sigma=1, seed=0):
- self.random_state = np.random.RandomState(seed)
- self.sigma = sigma
-
- def gaussian(self, T):
- # Get zero-mean unit variance gaussian distribution
- return self.sigma*self.random_state.randn(T)
-
- def weibull(self, T):
- # Get zero-mean sigma variance weibull distribution
- a = 2
- mean = scipy.special.gamma(1./a + 1)
- variance = scipy.special.gamma(
- 2./a + 1) - scipy.special.gamma(1./a + 1)**2
- return self.sigma*(self.random_state.weibull(a=a, size=T) - mean)/np.sqrt(variance)
-
- def uniform(self, T):
- # Get zero-mean sigma variance uniform distribution
- mean = 0.5
- variance = 1./12.
- return self.sigma*(self.random_state.uniform(size=T) - mean)/np.sqrt(variance)
-
-
-def lin_f(x): return x
-def f2(x): return (x + 5. * x**2 * np.exp(-x**2 / 20.))
-
-
-def simulate_random_var(seed, T, N, L, coef=[0.2, 0.8], auto_corr=[0.4,0.9], tau_max=5, noise_sigma=[0.01, 0.01]):
-
- if True:
- coupling_funcs = [lin_f]
- noise_types = ['gaussian'] # , 'weibull', 'uniform']
- # noise_sigma = (0.1, 0.3)
-
- couplings = list(np.arange(coef[0], coef[1]+1e-5, coef[2]))
- couplings += [-c for c in couplings]
-
- # auto_deps = list(np.arange(max(0., auto_corr-0.6), auto_corr+0.01, 0.05))
- auto_deps = list(np.arange(auto_corr[0], auto_corr[1]+1e-5, auto_corr[2]))
-
- # Models may be non-stationary. Hence, we iterate over a number of seeds
- # to find a stationary one regarding network topology, noises, etc
-
- ir = 0
- model_seed = seed
- while True:
- ir += 1
- # np.random.seed(model_seed)
- random_state = np.random.RandomState(model_seed)
-
- links = generate_random_contemp_model(
- N=N, L=L,
- coupling_coeffs=couplings,
- coupling_funcs=coupling_funcs,
- auto_coeffs=auto_deps,
- tau_max=tau_max,
- contemp_fraction=0.,
- # num_trials=1000,
- random_state=random_state)
-
- noises = []
- for j in links:
- noise_type = random_state.choice(noise_types)
- sigmas = list(np.arange(noise_sigma[0], noise_sigma[1]+1e-5, noise_sigma[2]))
- sigma = random_state.choice(sigmas)
- # sigma = noise_sigma[0] + (noise_sigma[1]-noise_sigma[0])*random_state.rand()
- noises.append(getattr(noise_model(sigma=sigma, seed=seed), noise_type))
-
- data_all_check, nonstationary = generate_nonlinear_contemp_timeseries(
- links=links, T=100, noises=noises, random_state=random_state)
-
- # If the model is stationary, break the loop
- if not nonstationary:
- data, nonstationary_full = generate_nonlinear_contemp_timeseries(
- links=links, T=T, noises=noises, random_state=random_state)
- if not nonstationary_full:
- break
- else:
- print("Trial %d: Not a stationary model" % ir)
- model_seed += 10000
-
- cm = links_to_matrix(links)
- return data, cm
-
-
-def simulate_var_from_links(links, T, seed=0, noise_sigma=[0.1, 0.2], noise_type="gaussian", func_name="lin_f"):
- """
- links_coeffs = {0: [((0, -1), 0.7), ((1, -1), -0.8)],
- 1: [((1, -1), 0.8), ((3, -1), 0.8)],
- 2: [((2, -1), 0.5), ((1, -2), 0.5), ((3, -3), 0.6)],
- 3: [((3, -1), 0.4)],
- }
- """
- def get_func(func_name):
- if func_name == "lin_f":
- return lin_f
- else:
- raise NotImplementedError
-
- random_state = np.random.RandomState(seed)
- noises = []
-
- new_links = {}
- for j in range(len(links)):
- sigma = noise_sigma[0] + \
- (noise_sigma[1]-noise_sigma[0])*random_state.rand()
- noises.append(getattr(noise_model(sigma=sigma, seed=seed), noise_type))
- new_links[j] = []
- for props in links[j]:
- new_links[j].append(
- (tuple(props[0:2]), props[2], get_func(props[3]),))
- data, nonstationary = generate_nonlinear_contemp_timeseries(
- links=new_links, T=T, noises=noises, random_state=random_state)
- if nonstationary:
- print("Model nonstationay!")
-
- cm = links_to_matrix(new_links)
- return data, cm
-
-
-def make_var_stationary(beta, radius=0.97):
- '''Rescale coefficients of VAR model to make stable.'''
- p = beta.shape[0]
- lag = beta.shape[1] // p
- bottom = np.hstack((np.eye(p * (lag - 1)), np.zeros((p * (lag - 1), p))))
- beta_tilde = np.vstack((beta, bottom))
- eigvals = np.linalg.eigvals(beta_tilde)
- max_eig = max(np.abs(eigvals))
- nonstationary = max_eig > radius
- if nonstationary:
- # print(f"Nonstationary, beta={str(beta):s}, max_eig={max_eig:.4f}")
- return make_var_stationary((beta / max_eig) * 0.7, radius)
- else:
- # print(f"Stationary, beta={str(beta):s}")
- return beta
-
-
-def simulate_var(p, T, lag, sparsity=0.2, beta_value=1.0, auto_corr=3.0, sd=0.1, seed=0):
- if seed is not None:
- np.random.seed(seed)
-
- # Set up coefficients and Granger causality ground truth.
- GC = np.eye(p, dtype=int)
- beta = np.eye(p) * auto_corr
-
- num_nonzero = int(p * sparsity) - 1
- for i in range(p):
- choice = np.random.choice(p - 1, size=num_nonzero, replace=False)
- choice[choice >= i] += 1
- beta[i, choice] = beta_value
- GC[i, choice] = 1
-
- beta = np.hstack([beta for _ in range(lag)])
- beta = make_var_stationary(beta)
-
- # Generate data.
- burn_in = 100
- errors = np.random.normal(loc=0, scale=sd, size=(p, T + burn_in))
- X = np.ones((p, T + burn_in))
- X[:, :lag] = errors[:, :lag]
- for t in range(lag, T + burn_in):
- X[:, t] = np.dot(beta, X[:, (t-lag):t].flatten(order='F'))
- X[:, t] += errors[:, t-1]
-
- data = X.T[burn_in:, :]
- return data, beta, GC
-
-
-
-
-
-def lorenz(x, t, F):
- '''Partial derivatives for Lorenz-96 ODE.'''
- p = len(x)
- dxdt = np.zeros(p)
- for i in range(p):
- dxdt[i] = (x[(i+1) % p] - x[(i-2) % p]) * x[(i-1) % p] - x[i] + F
-
- return dxdt
-
-
-def simulate_lorenz_96(p, T, F=10.0, delta_t=0.1, sd=0.1, burn_in=1000,
- seed=0):
- if seed is not None:
- np.random.seed(seed)
-
- # Use scipy to solve ODE.
- x0 = np.random.normal(scale=0.01, size=p)
- t = np.linspace(0, (T + burn_in) * delta_t, T + burn_in)
- X = odeint(lorenz, x0, t, args=(F,))
- X += np.random.normal(scale=sd, size=(T + burn_in, p))
-
- # Set up Granger causality ground truth.
- GC = np.zeros((p, p), dtype=int)
- for i in range(p):
- GC[i, i] = 1
- GC[i, (i + 1) % p] = 1
- GC[i, (i - 1) % p] = 1
- GC[i, (i - 2) % p] = 1
-
- return X[burn_in:, :], GC
-
diff --git a/singular_ticker_causal/algorithms/CUTS_PLUS/data/utils.py b/singular_ticker_causal/algorithms/CUTS_PLUS/data/utils.py
deleted file mode 100644
index c8b3c5297e256be7c92936960049ac6f09c80e40..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/algorithms/CUTS_PLUS/data/utils.py
+++ /dev/null
@@ -1,136 +0,0 @@
-import numpy as np
-import pandas as pd
-
-from sklearn.metrics.pairwise import haversine_distances
-
-
-def compute_mean(x, index=None):
- """Compute the mean values for each datetime. The mean is first computed hourly over the week of the year.
- Further NaN values are computed using hourly mean over the same month through the years. If other NaN are present,
- they are removed using the mean of the sole hours. Hoping reasonably that there is at least a non-NaN entry of the
- same hour of the NaN datetime in all the dataset."""
- if isinstance(x, np.ndarray) and index is not None:
- shape = x.shape
- x = x.reshape((shape[0], -1))
- df_mean = pd.DataFrame(x, index=index)
- else:
- df_mean = x.copy()
- cond0 = [df_mean.index.year, df_mean.index.isocalendar().week, df_mean.index.hour]
- cond1 = [df_mean.index.year, df_mean.index.month, df_mean.index.hour]
- conditions = [cond0, cond1, cond1[1:], cond1[2:]]
- while df_mean.isna().values.sum() and len(conditions):
- nan_mean = df_mean.groupby(conditions[0]).transform(np.nanmean)
- df_mean = df_mean.fillna(nan_mean)
- conditions = conditions[1:]
- if df_mean.isna().values.sum():
- df_mean = df_mean.fillna(method='ffill')
- df_mean = df_mean.fillna(method='bfill')
- if isinstance(x, np.ndarray):
- df_mean = df_mean.values.reshape(shape)
- return df_mean
-
-
-def geographical_distance(x=None, to_rad=True):
- """
- Compute the as-the-crow-flies distance between every pair of samples in `x`. The first dimension of each point is
- assumed to be the latitude, the second is the longitude. The inputs is assumed to be in degrees. If it is not the
- case, `to_rad` must be set to False. The dimension of the data must be 2.
-
- Parameters
- ----------
- x : pd.DataFrame or np.ndarray
- array_like structure of shape (n_samples_2, 2).
- to_rad : bool
- whether to convert inputs to radians (provided that they are in degrees).
-
- Returns
- -------
- distances :
- The distance between the points in kilometers.
- """
- _AVG_EARTH_RADIUS_KM = 6371.0088
-
- # Extract values of X if it is a DataFrame, else assume it is 2-dim array of lat-lon pairs
- latlon_pairs = x.values if isinstance(x, pd.DataFrame) else x
-
- # If the input values are in degrees, convert them in radians
- if to_rad:
- latlon_pairs = np.vectorize(np.radians)(latlon_pairs)
-
- distances = haversine_distances(latlon_pairs) * _AVG_EARTH_RADIUS_KM
-
- # Cast response
- if isinstance(x, pd.DataFrame):
- res = pd.DataFrame(distances, x.index, x.index)
- else:
- res = distances
-
- return res
-
-
-def infer_mask(df, infer_from='next'):
- """Infer evaluation mask from DataFrame. In the evaluation mask a value is 1 if it is present in the DataFrame and
- absent in the `infer_from` month.
-
- @param pd.DataFrame df: the DataFrame.
- @param str infer_from: denotes from which month the evaluation value must be inferred.
- Can be either `previous` or `next`.
- @return: pd.DataFrame eval_mask: the evaluation mask for the DataFrame
- """
- mask = (~df.isna()).astype('uint8')
- eval_mask = pd.DataFrame(index=mask.index, columns=mask.columns, data=0).astype('uint8')
- if infer_from == 'previous':
- offset = -1
- elif infer_from == 'next':
- offset = 1
- else:
- raise ValueError('infer_from can only be one of %s' % ['previous', 'next'])
- months = sorted(set(zip(mask.index.year, mask.index.month)))
- length = len(months)
- for i in range(length):
- j = (i + offset) % length
- year_i, month_i = months[i]
- year_j, month_j = months[j]
- mask_j = mask[(mask.index.year == year_j) & (mask.index.month == month_j)]
- mask_i = mask_j.shift(1, pd.DateOffset(months=12 * (year_i - year_j) + (month_i - month_j)))
- mask_i = mask_i[~mask_i.index.duplicated(keep='first')]
- mask_i = mask_i[np.in1d(mask_i.index, mask.index)]
- eval_mask.loc[mask_i.index] = ~mask_i.loc[mask_i.index] & mask.loc[mask_i.index]
- return eval_mask
-
-
-def disjoint_months(dataset, months=None):
- idxs = np.arange(len(dataset))
- months = ensure_list(months)
- # divide indices according to window or horizon
- start, end = 0, dataset.window - 1
- # after idxs
- start_in_months = np.in1d(dataset.index[dataset._indices + start].month, months)
- end_in_months = np.in1d(dataset.index[dataset._indices + end].month, months)
- idxs_in_months = start_in_months & end_in_months
- after_idxs = idxs[idxs_in_months]
- # previous idxs
- months = np.setdiff1d(np.arange(1, 13), months)
- start_in_months = np.in1d(dataset.index[dataset._indices + start].month, months)
- end_in_months = np.in1d(dataset.index[dataset._indices + end].month, months)
- idxs_in_months = start_in_months & end_in_months
- prev_idxs = idxs[idxs_in_months]
- return prev_idxs, after_idxs
-
-
-def thresholded_gaussian_kernel(x, theta=None, threshold=None, threshold_on_input=False):
- if theta is None:
- theta = np.std(x)
- weights = np.exp(-np.square(x / theta))
- if threshold is not None:
- mask = x > threshold if threshold_on_input else weights < threshold
- weights[mask] = 0.
- return weights
-
-
-def ensure_list(obj):
- if isinstance(obj, (list, tuple)):
- return list(obj)
- else:
- return [obj]
-
diff --git a/singular_ticker_causal/algorithms/CUTS_PLUS/lagged_graph.py b/singular_ticker_causal/algorithms/CUTS_PLUS/lagged_graph.py
deleted file mode 100644
index e6dd4bae8311a59fa32ce94794c52bf3428dd526..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/algorithms/CUTS_PLUS/lagged_graph.py
+++ /dev/null
@@ -1,95 +0,0 @@
-from __future__ import annotations
-
-from copy import deepcopy
-from dataclasses import dataclass
-from typing import Any
-
-import numpy as np
-
-from .cuts_plus import main as cuts_plus_main
-
-
-@dataclass(frozen=True)
-class LaggedGraphBundle:
- graphs: np.ndarray
- summary_graph: np.ndarray
- priors: np.ndarray
- input_steps: np.ndarray
-
-
-def broadcast_prior_to_lags(G_prior: np.ndarray | None, num_lags: int) -> np.ndarray | None:
- if G_prior is None:
- return None
- G_prior = np.asarray(G_prior, dtype=np.float32)
- if G_prior.ndim == 2:
- return np.repeat(G_prior[None, :, :], num_lags, axis=0)
- if G_prior.ndim == 3 and G_prior.shape[0] == num_lags:
- return G_prior
- raise ValueError("G_prior must have shape (N, N) or (L, N, N)")
-
-
-def aggregate_lagged_graphs(graphs: np.ndarray, reducer: str = "max") -> np.ndarray:
- if reducer == "max":
- return np.max(graphs, axis=0)
- if reducer == "mean":
- return np.mean(graphs, axis=0)
- if reducer == "last":
- return graphs[-1]
- raise ValueError(f"Unsupported lagged graph reducer: {reducer}")
-
-
-def discover_lagged_graphs(
- data: np.ndarray,
- mask: np.ndarray,
- opt: Any,
- log: Any,
- device: str = "cpu",
- text_data: np.ndarray | None = None,
- text_mask: np.ndarray | None = None,
- G_prior: np.ndarray | None = None,
- num_lags: int = 1,
- reducer: str = "max",
-) -> LaggedGraphBundle:
- if num_lags < 1:
- raise ValueError("num_lags must be >= 1")
-
- lag_priors = broadcast_prior_to_lags(G_prior, num_lags)
- graphs = []
- input_steps = []
- previous_graph = None
-
- for lag_idx in range(num_lags):
- cfg = deepcopy(opt)
- cfg.input_step = max(int(opt.input_step), lag_idx + 1)
- input_steps.append(cfg.input_step)
-
- current_prior = None
- if lag_priors is not None:
- current_prior = lag_priors[lag_idx]
- elif previous_graph is not None:
- current_prior = previous_graph
-
- graph = cuts_plus_main(
- data=data,
- mask=mask,
- true_cm=None,
- opt=cfg,
- log=log,
- device=device,
- text_data=text_data,
- text_mask=text_mask,
- G_prior=current_prior,
- )
- previous_graph = graph
- graphs.append(np.asarray(graph, dtype=np.float32))
-
- stacked = np.stack(graphs, axis=0)
- summary = aggregate_lagged_graphs(stacked, reducer=reducer)
- if lag_priors is None:
- lag_priors = np.zeros_like(stacked)
- return LaggedGraphBundle(
- graphs=stacked,
- summary_graph=summary.astype(np.float32),
- priors=lag_priors.astype(np.float32),
- input_steps=np.asarray(input_steps, dtype=np.int32),
- )
diff --git a/singular_ticker_causal/algorithms/CUTS_PLUS/model/cuts_plus_net.py b/singular_ticker_causal/algorithms/CUTS_PLUS/model/cuts_plus_net.py
deleted file mode 100644
index 8bd00bb6dd85b64faf2ede6df6ae1a51ee5717d2..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/algorithms/CUTS_PLUS/model/cuts_plus_net.py
+++ /dev/null
@@ -1,153 +0,0 @@
-import torch
-from einops import rearrange
-from torch import nn
-
-class GRUCell(nn.Module):
-
- def __init__(self, d_in, num_units, n_nodes, concat_h=False, activation='tanh'):
- super(GRUCell, self).__init__()
- self.activation_fn = getattr(torch, activation)
-
- mpnn_channel = d_in*n_nodes+num_units if concat_h else d_in*n_nodes
- self.forget_gate = MPNN(c_in=mpnn_channel, c_out=num_units, concat_h=concat_h)
- self.update_gate = MPNN(c_in=mpnn_channel, c_out=num_units, concat_h=concat_h)
- self.c_gate = MPNN(c_in=mpnn_channel, c_out=num_units, concat_h=concat_h)
-
- def forward(self, x, h, adj):
- """
- :param x: (B, input_dim, num_nodes)
- :param h: (B, num_units, num_nodes)
- :param adj: (num_nodes, num_nodes)
- :return:
- """
- # we start with bias 1.0 to not reset and not update
- r = torch.sigmoid(self.forget_gate(x, h, adj))
- u = torch.sigmoid(self.update_gate(x, h, adj))
- c = self.c_gate(x, r * h, adj) # batch_size, self._num_nodes * output_size
- c = self.activation_fn(c)
- return u * h + (1. - u) * c
-
-
-class MPNN(nn.Module):
- def __init__(self, c_in, c_out, concat_h=True):
- super(MPNN, self).__init__()
- self.concat_h = concat_h
- self.mlp = nn.Conv1d(c_in, c_out, kernel_size=1)
-
- def forward(self, x, h, graph):
- b, c, n = x.shape
-
- x_repeat = x[:, :, :, None].expand(-1, -1, -1, n) # [b, c, n, n]
- # graph = rearrange(graph, 'b n m -> b m n')
- x_messages = torch.einsum('bcmn,bmn->bcmn', (x_repeat, graph))
- x_messages = rearrange(x_messages, 'b c m n -> b (c m) n')
-
- if self.concat_h:
- out = self.mlp(torch.cat([x_messages, h], dim=1))
- else:
- out = self.mlp(x_messages)
- return out
-
-
-class LocalConv1D(nn.Module):
- def __init__(self, in_channels, out_channels, kernel_size, n_nodes):
- super(LocalConv1D, self).__init__()
- self.out_channel = out_channels
- self.conv_list = nn.ModuleList([
- nn.Conv1d(in_channels=in_channels, out_channels=out_channels, kernel_size=kernel_size) for _ in range(n_nodes)
- ])
-
- def forward(self, x): # x: [batch, features, nodes]
- b, h, n = x.shape
- out = torch.zeros((b, self.out_channel, n)).to(x.device)
- for i in range(n):
- x_local_in = x[..., i].unsqueeze(-1)
- x_local_out = self.conv_list[i](x_local_in)
- out[..., i] = x_local_out.squeeze(-1)
- return out
-
-
-class CUTS_Plus_Net(nn.Module):
- def __init__(self, n_nodes,
- in_ch=1,
- hidden_ch=32,
- n_layers=1,
- shared_weights_decoder=False,
- concat_h=False,):
- super().__init__()
- self.in_ch = in_ch
- self.hidden_ch = hidden_ch
- self.n_layers = n_layers
-
- self.conv_encoder1 = nn.Conv1d(in_channels=hidden_ch, out_channels=hidden_ch, kernel_size=1)
- self.conv_encoder2 = nn.Conv1d(in_channels=2*hidden_ch, out_channels=hidden_ch, kernel_size=1)
- if shared_weights_decoder:
- self.decoder = nn.Sequential(
- nn.Conv1d(in_channels=2*hidden_ch, out_channels=in_ch, kernel_size=1),
- # nn.LeakyReLU(),
- # nn.Conv1d(in_channels=hidden_ch, out_channels=hidden_ch, kernel_size=1),
- # nn.LeakyReLU(),
- # nn.Conv1d(in_channels=hidden_ch, out_channels=in_ch, kernel_size=1),
- # nn.LeakyReLU(),
- )
- else:
- self.decoder = nn.Sequential(
- LocalConv1D(in_channels=2*hidden_ch, out_channels=in_ch, kernel_size=1, n_nodes=n_nodes),
- # nn.LeakyReLU(),
- # LocalConv1D(in_channels=hidden_ch, out_channels=hidden_ch, kernel_size=1, n_nodes=n_nodes),
- # nn.LeakyReLU(),
- # LocalConv1D(in_channels=hidden_ch, out_channels=in_ch, kernel_size=1, n_nodes=n_nodes),
- # nn.LeakyReLU(),
- )
- # self.act = nn.PReLU()
- self.act = nn.LeakyReLU()
-
- self.cells = nn.ModuleList()
- for i in range(self.n_layers):
- self.cells.append(GRUCell(d_in=in_ch if i==0 else hidden_ch,
- num_units=hidden_ch,
- n_nodes=n_nodes,
- concat_h=concat_h))
-
- self.h0 = self.init_state(n_nodes)
-
- def init_state(self, n_nodes):
- h = []
- for layer in range(self.n_layers):
- h.append(nn.parameter.Parameter(torch.zeros([self.hidden_ch, n_nodes])))
- return nn.ParameterList(h)
-
- def update_state(self, x, h, graph):
- rnn_in = x
- for layer in range(self.n_layers):
- rnn_in = h[layer] = self.cells[layer](rnn_in, h[layer], graph)
- return h
-
- def forward(self, x, mask, fwd_graph):
- x = rearrange(x, 'b n s c -> b c n s')
- # fwd_graph = torch.ones_like(fwd_graph)
- # mask = torch.ones_like(x).byte()
- bs, in_ch, n_nodes, steps = x.shape
-
- h = [h_.expand(bs, -1, -1) for h_ in self.h0.to(x.device)]
-
- pred = []
- for step in range(steps):
- x_now = x[..., step] # [batches, in_ch, nodes]
-
- """Update state"""
- h = self.update_state(x_now, h, fwd_graph)
- h_now = h[-1]
-
- """Prediction"""
- x_repr = self.act(self.conv_encoder1(h_now)) # [batches, hidden_ch, nodes]
- x_repr = self.act(self.conv_encoder2(torch.cat([x_repr, h_now], dim=1))) # [batches, hidden_ch, nodes]
- x_repr = torch.cat([x_repr, h_now], dim=1) # [batches, 2*hidden_ch, nodes]
- x_hat2 = self.decoder(x_repr) # [batches, in_ch, nodes]
- pred.append(x_hat2)
-
-
- pred = torch.stack(pred, dim=-1)
- pred = rearrange(pred, 'b c n s -> b n s c')
- return pred[:, :, -1:]
-
diff --git a/singular_ticker_causal/algorithms/CUTS_PLUS/utils/causal_plot.py b/singular_ticker_causal/algorithms/CUTS_PLUS/utils/causal_plot.py
deleted file mode 100644
index e8d98081b197aaf76a81596f891e9df68249db1a..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/algorithms/CUTS_PLUS/utils/causal_plot.py
+++ /dev/null
@@ -1,57 +0,0 @@
-from re import X
-import matplotlib
-matplotlib.use('Agg')
-import matplotlib.pyplot as plt
-import tqdm
-import numpy as np
-import matplotlib.cm as cm
-
-def save_causal_graph(save_path, causal_matrix: np.ndarray, thres_percentile=100, colormap="gnuplot2"):
- causal_matrix = np.max(causal_matrix)
- print(causal_matrix.shape)
-
- n_node = causal_matrix.shape[0]
- image_size = [100, 100]
- n_node_dim = n_node**0.5
- causal_thres = np.percentile(causal_matrix, 100-thres_percentile)
-
- colormap = cm.get_cmap(colormap)
-
- plt.figure(figsize=[10,10], facecolor='black', edgecolor='black')
-
-
- for node_i_from in tqdm.tqdm(range(n_node)):
- x_from = image_size[0] // n_node_dim * (node_i_from // n_node_dim + 0.5)
- y_from = image_size[0] // n_node_dim * (node_i_from % n_node_dim + 0.5)
- plt.text(y=x_from, x=y_from, s=f"{node_i_from:d}", size=20, color="#ffffff")
- for node_i_to in range(n_node):
- if not node_i_from == node_i_to:
- x_to = image_size[0] // n_node_dim * (node_i_to // n_node_dim + 0.5)
- y_to = image_size[0] // n_node_dim * (node_i_to % n_node_dim + 0.5)
- causal_effect = causal_matrix[node_i_from, node_i_to]
- if causal_effect > causal_thres:
- width = max(0.01, 1*causal_effect)
- arrow_length = ((x_to-x_from)**2 + (y_to-y_from)**2)**0.5
- plt.arrow(
- y=x_from+(x_to-x_from)*width/arrow_length,
- x=y_from+(y_to-y_from)*width/arrow_length,
- dy=(x_to-x_from)*(arrow_length-5*width)/arrow_length,
- dx=(y_to-y_from)*(arrow_length-5*width)/arrow_length,
- width=width,
- head_length=4*width,
- facecolor=colormap(causal_effect)[:3]+(causal_effect,),
- edgecolor="#00000000"
- )
-
- # fig.add_annotation(text=f"{node_i_from:d}", x=x_from, y=y_from, showarrow=False
-
- ax=plt.gca()
- ax.patch.set_facecolor("black")
- ax.xaxis.set_ticks_position('top')
- ax.invert_yaxis()
- plt.savefig(save_path)
-
-
-if __name__=="__main__":
- causal_matrix = np.load("outputs/tsgae_2022_0716_191203_262072/w.npy")
- save_causal_graph("outputs/pic/causal.png", causal_matrix, thres_percentile=100)
\ No newline at end of file
diff --git a/singular_ticker_causal/algorithms/CUTS_PLUS/utils/exp_utils.py b/singular_ticker_causal/algorithms/CUTS_PLUS/utils/exp_utils.py
deleted file mode 100644
index 167bb5e7001d3ad104042d305c02ff4ac403cdd6..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/algorithms/CUTS_PLUS/utils/exp_utils.py
+++ /dev/null
@@ -1,286 +0,0 @@
-import os
-from os.path import join as opj
-from os.path import dirname as opd
-
-from .opt_type import MultiCADopt
-from .misc import omegaconf2dict
-
-import re
-import numpy as np
-import matplotlib
-matplotlib.use('Agg')
-import matplotlib.pyplot as plt
-
-from datetime import datetime
-from omegaconf import OmegaConf
-import glob
-import tqdm
-from tensorboard.backend.event_processing import event_accumulator
-
-PROPER_NAME = {"NFGR":"BRIEF", "h265":"H.265", "h264":"H.264", "jpg":"JPEG", "aoi-2000":"AoI", "vvc":"H.266"}
-
-MAIN_COLOR = {"NFGR":"#e64b35", "h265":"#48c9b0", "h264":"#5599c7", "jpg":"#c39bd2", "DVC":"#e6b0aa", "SGA":"#f8c370", "vvc":"#b1babb",
- "GOOD":"#239954", "BAD":"#cc6155"}
-SECONDARY_COLOR = {"NFGR":"#d98880", "h265":"#76d7c3", "SGA":"#f9d7a0", "h264":"#7fb3d5", "DVC":"#e6b0aa"}
-
-
-def get_time_stamp():
- return str(datetime.now().strftime("%m-%d-%H%M%S-%f"))
-
-def show_single_scores(x_arr, y_arr, label="exp", suffix="",
- scatter=False, log_axis=True,
- xlim=None, ylim=None):
-
- '''Plotting figures'''
- cvs = FigureCanvas(name=label, figsize=[30,20])
- fig_idx = 0
- for idx,data_path in enumerate(y_arr.labels[2]): # range(score_arr.shape[2]):
- if fig_idx == 40:
- break
- elif np.isnan(np.nanmean(y_arr[:,:,data_path])):
- continue
-
- fig_idx += 1
- if log_axis:
- ax = plt.subplot(8,5,fig_idx, xscale="log")
- else:
- ax = plt.subplot(8,5,fig_idx)
- ax.set_title("DATA_{:02d}_".format(idx) + data_path[-100:-60] + "\n" + data_path[-60:])
- # ax.set_title(dim_marks[2][data_i])
- plt.set_cmap("rainbow")
- for i, dim0 in enumerate(y_arr.labels[0]):
- x_nan = x_arr[dim0,:,data_path]
- y_nan = y_arr[dim0,:,data_path]
- x = x_nan[np.isfinite(x_nan + y_nan)]
- y = y_nan[np.isfinite(x_nan + y_nan)]
-
- if len(x) > 0 and len(y) > 0:
- x, y = sort_lists(x, y)
-
- plt.plot(x, y, color=plt.get_cmap("tab20")(i), label=dim0)
- plt.scatter(x, y, color=plt.get_cmap("tab20")(i))
-
- plt.legend()
- if ylim is not None:
- plt.ylim(ylim)
-
- ax=plt.gca()
- ax.xaxis.set_major_locator(plt.LogLocator(base=10, numticks=5))
- ax.yaxis.set_major_locator(plt.MaxNLocator(5))
-
- cvs.save_fig(suffix=suffix, time_stamp=False, save_format=".pdf")
-
-
-def sort_lists(*lists):
- sorted_index = np.argsort(lists[0]).astype(int)
- results = []
- for l in lists:
- sorted_list = np.array([l[i] for i in sorted_index])
- results.append(sorted_list)
- return results
-
-
-def show_averge_scores(x_arr, y_arr, label="exp", suffix="", percentile=25,
- scatter=False, log_axis=True, figsize=[4,3],
- xlim=None, ylim=None, legend=False, grid=False, std=False):
- cvs = FigureCanvas(name=label, figsize=figsize)
- if log_axis:
- plt.xscale("log")
-
- fig_idx = 0
- if scatter:
- for data_i,data_path in enumerate(y_arr.labels[2]):
- if np.isnan(np.nanmean(y_arr[:,:,data_path])):
- continue
-
- for i, dim0 in enumerate(y_arr.labels[0]):
- x_nan = x_arr[dim0,:,data_path]
- y_nan = y_arr[dim0,:,data_path]
- x = x_nan[np.isfinite(x_nan + y_nan)]
- y = y_nan[np.isfinite(x_nan + y_nan)]
-
- if len(x) > 0 and len(y) > 0:
- x, y = sort_lists(x, y)
-
- # plt.plot(x, y, color=plt.get_cmap("tab20")(i), label=dim0, marker="v")
- plt.scatter(x, y, color=MAIN_COLOR[dim0], alpha=0.5, marker="v", edgecolors='none',
- s=80 if "NFGR" in dim0 else 50)
-
- # plt.legend(edgecolors='none')
- if ylim is not None:
- plt.ylim(ylim)
-
- if ylim is not None:
- full_range = ylim[1] - ylim[0]
- else:
- full_range = np.nanmax(np.nanmean(y_arr.arr, axis=2)) - np.nanmin(np.nanmean(y_arr.arr, axis=2))
- max_std = np.nanmax(np.nanstd(y_arr["NFGR",:,:], axis=1))
-
-
- for idx, dim0 in enumerate(sorted(y_arr.labels[0], key=lambda item:item == "NFGR")):
- # if "jpg" in dim0:
- # continue
-
- x_nan = np.nanmean(x_arr[dim0,:], axis=1)
- y_nan = np.nanmean(y_arr[dim0,:], axis=1)
- # y_l = np.nanpercentile(y_arr[dim0,:], percentile, axis=1)
- # y_u = np.nanpercentile(y_arr[dim0,:], 100-percentile, axis=1)
- y_std = np.nanstd(y_arr[dim0,:], axis=1) # / max_std * full_range * 0.07
-
- print(1 / max_std * full_range * 0.07)
-
- x = x_nan[np.isfinite(x_nan + y_nan)]
- y = y_nan[np.isfinite(x_nan + y_nan)]
- # y_l = y_l[np.isfinite(x_nan + y_nan)]
- # y_u = y_u[np.isfinite(x_nan + y_nan)]
- y_std = y_std[np.isfinite(x_nan + y_nan)]
-
- x, y, y_std = sort_lists(x, y, y_std)
-
- if std:
- plt.fill_between(x, y-y_std/2, y+y_std/2, color=MAIN_COLOR[dim0], alpha=0.25, edgecolors="none")
- plt.plot(x, y, color=MAIN_COLOR[dim0], label=name(dim0),
- lw=2 if "NFGR" in dim0 else 1.5)
- plt.scatter(x, y, color=MAIN_COLOR[dim0],
- s=30 if "NFGR" in dim0 else 20)
-
- if ylim is not None:
- plt.ylim(ylim)
- if xlim is not None:
- plt.xlim(xlim)
- if legend:
- plt.legend(loc='lower left', bbox_to_anchor=(0.1, 0.1), fancybox=False)
-
- ax=plt.gca()
- if log_axis:
- ax.xaxis.set_major_locator(plt.LogLocator(base=10, numticks=5))
- else:
- ax.xaxis.set_major_locator(plt.MaxNLocator(5))
- ax.yaxis.set_major_locator(plt.MaxNLocator(5))
-
- if grid:
- ax.spines['right'].set_visible(True)
- ax.spines['top'].set_visible(True)
- # plt.xticks(np.arange(0.4, 1.8, 0.28))
- # plt.yticks(np.arange(50, 500, 40))
- plt.grid(axis='both', c="#cacaca", which="major")
-
- cvs.save_fig(suffix=suffix, time_stamp=False, save_format=".pdf")
-
-
-def name(alias):
- if alias in PROPER_NAME:
- return PROPER_NAME[alias]
- else:
- print("Cannot find proper name.")
- return alias
-
-class FigureCanvas(object):
-
- def __init__(self, name="ex1", figsize=[14,9]):
- self.name = name
- plt.close('all')
- fig = plt.figure(figsize=figsize)
- ax = plt.axes()
- ax.spines['right'].set_visible(False)
- ax.spines['top'].set_visible(False)
- plt.tight_layout()
-
- def show_fig(self, save_format=".png", suffix="", time_stamp=True):
- save_path = "./exp/figs/%s/%s_%s%s"%(
- self.name,
- get_time_stamp() if time_stamp else "plt",
- suffix, save_format)
- if not os.path.exists(opd(save_path)):
- os.makedirs(opd(save_path))
- plt.savefig(save_path, bbox_inches='tight')
- plt.show()
-
- def save_fig(self, save_format=".png", suffix="", time_stamp=True, save_root="./exp/figs/"):
- save_path = opj(save_root, "%s/%s_%s%s"%(
- self.name,
- get_time_stamp() if time_stamp else "plt",
- suffix, save_format))
- if not os.path.exists(opd(save_path)):
- os.makedirs(opd(save_path))
- plt.savefig(save_path, bbox_inches='tight')
-
-
-def find_lineprofile_cmp(im_list):
- for x in range(0, im_list[0].shape[0], 10):
- for y in range(0, im_list[0].shape[1], 10):
- lp = [im[x][y] for im in im_list]
- if np.max(lp) > 2500 and np.max(lp) < 3000:
- return lp
-
-
-def get_decompressed_path(opt_path):
- res_root = opd(opt_path)
- max_step = 0
- for dirn in os.listdir(res_root):
- if dirn == "decompressed":
- return glob.glob(res_root + "/decompressed/*.tif")[0]
- elif "steps" in dirn:
- step_n = int(dirn[5:])
- if step_n > max_step:
- max_step = step_n
- search_list = glob.glob(res_root + "/steps" + str(max_step) + "/decompressed/*.*")
- if len(search_list) > 0:
- return search_list[0]
- return None
-
-def load_scalars(event_path):
- try:
- event_path = glob.glob(event_path)[0]
- except:
- print("No event file found.")
- return None
- ea = event_accumulator.EventAccumulator(event_path)
- ea.Reload()
- # print("Available scalars: ", ea.scalars.Keys())
- scalars = {}
- for criterion in ea.scalars.Keys():
- val_scalar = ea.scalars.Items(criterion)
- val_curve = ([(i.step, i.value) for i in val_scalar])
- scalars[criterion] = val_curve
- return scalars
-
-def load_scalars_cached(root_path, cache_dir="exp/cache", reload_data=False):
-
- if root_path[-1] == "/":
- root_path = root_path[:-1]
-
- # csv_file = glob.glob(root_path + "/*.csv")[0]
- # exp_list = load_csv(csv_file)
- root_name = "".join(re.split("/|\\\\", root_path)[-2:])
- cache_path = opj(cache_dir, root_name + ".npy")
-
- # read from cached data
- if os.path.exists(cache_path) and not reload_data:
- print("Loading cached result...")
- loaded_res = np.load(cache_path, allow_pickle=True)
- else:
- if not os.path.exists(cache_dir):
- os.makedirs(cache_dir)
-
- loaded_res = []
- for dirn in tqdm.tqdm(os.listdir(root_path)):
- if os.path.isdir(opj(root_path, dirn)):
- if len(glob.glob(root_path + "/%s/events.out.tfevents*"%(dirn))) > 0:
- res_fname = glob.glob(root_path + "/%s/events.out.tfevents*"%(dirn))[0]
- opt_fname = glob.glob(root_path + "/%s/opt.yaml"%(dirn))[0]
- scores = load_scalars(res_fname)
- opt: MultiCADopt = omegaconf2dict(OmegaConf.load(opt_fname), sep=".")
- loaded_res.append((res_fname, scores, opt))
-
- np.save(cache_path, loaded_res)
-
- return loaded_res
-
-
-# if __name__=="__main__":
-# # scalars = load_scalars("cyx_exp/experiments_outputs/ex2_1227/*/exp_00000/events.out.tfevents*")
-# # print(scalars)
-
-# res = load_scalars_cached("cyx_exp/experiments_outputs/ex2_1227/", reload_data=False)
-# print("")
\ No newline at end of file
diff --git a/singular_ticker_causal/algorithms/CUTS_PLUS/utils/gumbel_softmax.py b/singular_ticker_causal/algorithms/CUTS_PLUS/utils/gumbel_softmax.py
deleted file mode 100644
index 513e66a31ef978a2396213e51a6236bb934f5091..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/algorithms/CUTS_PLUS/utils/gumbel_softmax.py
+++ /dev/null
@@ -1,66 +0,0 @@
-import torch
-from torch import Tensor
-import warnings
-
-def gumbel_softmax(logits: Tensor, tau: float = 1, hard: bool = False, eps: float = 1e-10, dim: int = -1) -> Tensor:
- r"""
- Samples from the Gumbel-Softmax distribution (`Link 1`_ `Link 2`_) and optionally discretizes.
-
- Args:
- logits: `[..., num_features]` unnormalized log probabilities
- tau: non-negative scalar temperature
- hard: if ``True``, the returned samples will be discretized as one-hot vectors,
- but will be differentiated as if it is the soft sample in autograd
- dim (int): A dimension along which softmax will be computed. Default: -1.
-
- Returns:
- Sampled tensor of same shape as `logits` from the Gumbel-Softmax distribution.
- If ``hard=True``, the returned samples will be one-hot, otherwise they will
- be probability distributions that sum to 1 across `dim`.
-
- .. note::
- This function is here for legacy reasons, may be removed from nn.Functional in the future.
-
- .. note::
- The main trick for `hard` is to do `y_hard - y_soft.detach() + y_soft`
-
- It achieves two things:
- - makes the output value exactly one-hot
- (since we add then subtract y_soft value)
- - makes the gradient equal to y_soft gradient
- (since we strip all other gradients)
-
- Examples::
- >>> logits = torch.randn(20, 32)
- >>> # Sample soft categorical using reparametrization trick:
- >>> F.gumbel_softmax(logits, tau=1, hard=False)
- >>> # Sample hard categorical using "Straight-through" trick:
- >>> F.gumbel_softmax(logits, tau=1, hard=True)
-
- .. _Link 1:
- https://arxiv.org/abs/1611.00712
- .. _Link 2:
- https://arxiv.org/abs/1611.01144
- """
- if eps != 1e-10:
- warnings.warn("`eps` parameter is deprecated and has no effect.")
-
- gumbels = (
- -torch.empty_like(logits, memory_format=torch.legacy_contiguous_format).exponential_().log()
- ) # ~Gumbel(0,1)
- gumbels = (logits + gumbels) / tau # ~Gumbel(logits,tau)
- y_soft = gumbels.softmax(dim)
-
- if hard:
- # Straight through.
- index = y_soft.max(dim, keepdim=True)[1]
- y_hard = torch.zeros_like(logits, memory_format=torch.legacy_contiguous_format).scatter_(dim, index, 1.0)
- ret = y_hard - y_soft.detach() + y_soft
- else:
- # Reparametrization trick.
- ret = y_soft
- return ret
-
-if __name__=="__main__":
- a = torch.tensor([[2.0, 0.7]]*10)
- print(gumbel_softmax(a, tau=10))
\ No newline at end of file
diff --git a/singular_ticker_causal/algorithms/CUTS_PLUS/utils/logger.py b/singular_ticker_causal/algorithms/CUTS_PLUS/utils/logger.py
deleted file mode 100644
index e2525d49687c632bc2a5df2e035f941aab91dddc..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/algorithms/CUTS_PLUS/utils/logger.py
+++ /dev/null
@@ -1,74 +0,0 @@
-import sys
-from omegaconf import OmegaConf
-import os
-from os.path import join as opj
-import numpy as np
-from os.path import dirname as opd
-from typing import Dict
-from torch.utils.tensorboard import SummaryWriter
-from .misc import omegaconf2list
-
-
-class MyLogger():
- def __init__(self, log_dir: str, stderr: bool = True, tensorboard: bool = True, stdout: bool = True):
- self.log_dir = log_dir
- if not os.path.exists(self.log_dir):
- os.makedirs(self.log_dir)
- self.logger_dict: Dict[str] = {}
- if stdout:
- stdout_handler = open(opj(self.log_dir, 'stdout.log'), 'w')
- sys.stdout = stdout_handler
- if stderr:
- stderr_handler = open(opj(self.log_dir, 'stderr.log'), 'w')
- sys.stderr = stderr_handler
- if tensorboard:
- self.tblogger = SummaryWriter(self.log_dir)
- self.logger_dict['tblogger'] = self.tblogger
-
- def log_opt(self, opt):
- OmegaConf.save(config=opt, f=opj(self.log_dir, 'opt.yaml'))
- opt_log = omegaconf2list(opt, sep='/')
- for logger_name in self.logger_dict.keys():
- if logger_name == 'tblogger':
- for idx, opt in enumerate(opt_log):
- self.logger_dict[logger_name].add_text('hparam', opt, idx)
-
- def log_metrics(self, metrics_dict: Dict[str, float], iters):
- for logger_name in self.logger_dict.keys():
- if logger_name == 'csvlogger':
- self.logger_dict[logger_name].log_metrics(metrics_dict, iters)
- self.logger_dict[logger_name].save()
- elif logger_name == 'clearml_logger':
- for k in metrics_dict.keys():
- self.logger_dict[logger_name].report_scalar(
- k, k, metrics_dict[k], iters)
- elif logger_name == 'tblogger':
- for k in metrics_dict.keys():
- self.logger_dict[logger_name].add_scalar(
- k, metrics_dict[k], iters)
-
- def log_figures(self, figure, name="figure.png", iters=None, exclude_logger=[]):
- for logger_name in self.logger_dict.keys():
- if logger_name == 'tblogger':
- if logger_name not in exclude_logger:
- self.logger_dict[logger_name].add_figure(tag=name, figure=figure, global_step=iters)
-
- if iters is None:
- save_path = opj(self.log_dir, "figures")
- else:
- save_path = opj(self.log_dir, f"iter_{iters:d}", name)
- os.makedirs(opd(save_path), exist_ok=True)
- figure.savefig(save_path)
-
- def log_npz(self, data: Dict, name="data.npz", iters=None):
- if iters is None:
- save_path = opj(self.log_dir)
- else:
- save_path = opj(self.log_dir, f"iter_{iters:d}", name)
- os.makedirs(save_path, exist_ok=True)
- np.savez(opj(save_path, "graph.npz"), **data)
-
- def close(self):
- for logger_name in self.logger_dict.keys():
- if logger_name == 'tblogger':
- self.logger_dict[logger_name].close()
diff --git a/singular_ticker_causal/algorithms/CUTS_PLUS/utils/misc.py b/singular_ticker_causal/algorithms/CUTS_PLUS/utils/misc.py
deleted file mode 100644
index 9a686e29e2757c67cc53128c422c1099551acf5d..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/algorithms/CUTS_PLUS/utils/misc.py
+++ /dev/null
@@ -1,292 +0,0 @@
-from copy import deepcopy
-import numpy as np
-import matplotlib
-matplotlib.use('Agg')
-import matplotlib.pyplot as plt
-import itertools
-import random
-import numpy as np
-import torch
-import omegaconf
-from sklearn.metrics import roc_curve, roc_auc_score
-
-
-def log_time_series(original_data, data_interp, data_pred, log, log_step):
- fig = plt.figure(figsize=[10,10])
- plt.plot(np.arange(0, original_data.shape[0], 1), original_data, label="original")
- plt.plot(np.arange(0, data_interp.shape[0], 1), data_interp, label="interp")
- plt.plot(np.arange(0, data_pred.shape[0], 1), data_pred, label="pred")
- plt.legend()
- log.log_figures(fig, name="Predicted Latent Data", iters=log_step)
-
-
-def calc_and_log_metrics(time_prob_mat, true_cm, log, log_step, threshold=0.5, plot_roc=True):
- if len(time_prob_mat.shape) == 3:
- graph = np.max(time_prob_mat, axis=2)
- else:
- graph = time_prob_mat
- causal_graph = graph > threshold
- tp = np.mean(causal_graph * true_cm)
- tn = np.mean((1-causal_graph) * (1-causal_graph))
- fp = np.mean(causal_graph * (1-true_cm))
- fn = np.mean((1-causal_graph) * true_cm)
- tpr = tp / (tp + fn)
- fpr = fp / (fp + tn)
- acc = (tp + tn) / (tp + tn + fp + fn)
- log.log_metrics({"metrics/tpr": tpr}, log_step)
- log.log_metrics({"metrics/fpr": fpr}, log_step)
- log.log_metrics({"metrics/accuracy": acc}, log_step)
-
- if plot_roc:
- fpr, tpr, thres = roc_curve(true_cm.reshape(-1) > 0.5,
- graph.reshape(-1), pos_label=1)
- fig = plt.figure(figsize=[4, 4])
- plt.plot(fpr, tpr)
- log.tblogger.add_figure(tag="ROC", figure=fig, global_step=log_step)
-
- log.log_npz(name="graph",
- data={"true_cm":true_cm, "pred_cm":graph},
- iters=log_step)
-
- auc = roc_auc_score(true_cm.reshape(-1)>0.5,
- graph.reshape(-1))
- log.log_metrics({"metrics/auc": auc}, log_step)
- return auc
-
-def sigmoid(z):
- return 1/(1 + np.exp(-z))
-
-def plot_causal_matrix_in_training(time_coef, name, log, log_step, threshold=0.5, plot_each_time=False):
- if time_coef is None:
- return
-
- if np.max(time_coef) - np.min(time_coef) > 0.01:
- time_coef = (time_coef - np.min(time_coef)) / (np.max(time_coef) - np.min(time_coef))
- n, m, t = time_coef.shape
-
- # # Show Discovered Graph (Coefficiency)
- # sub_cg = plot_causal_matrix(
- # np.max(time_coef, axis=2),
- # figsize=[1.5*time_coef.shape[0], 1*n])
- # log.log_figures(sub_cg, name="Discovered Graph Coef/" + name, iters=log_step)
-
- # # Graph for Each Time Lag
- # if plot_each_time:
- # for ti in range(t):
- # sub_cg = plot_causal_matrix(
- # time_coef[:, :, ti],
- # figsize=[1.5*n, 1*n],
- # vmin=0, vmax=1)
- # log.log_figures(sub_cg, name=f"Discovered Prob T-{t-ti:d}",
- # iters=log_step, exclude_logger="tblogger")
-
- # Show Discovered Graph (Probability)
- time_graph = time_coef
- sub_cg = plot_causal_matrix(
- np.max(time_graph, axis=2),
- figsize=[1.5*n, 1*n],
- vmin=0, vmax=1)
- log.log_figures(sub_cg, name="Discovered Prob/" + name, iters=log_step)
-
- # Show Thresholded Graph
- time_thres = np.max(time_graph, axis=2) > threshold
- sub_cg = plot_causal_matrix(
- time_thres,
- figsize=[1.5*n, 1*n])
- log.log_figures(sub_cg, name="Discovered Graph/" + name, iters=log_step)
- log.log_npz({"Discovered Graph Coef": time_coef, "Discovered Prob": time_graph, "Discovered Graph": time_thres},
- name="Graph.npz", iters=log_step)
-
-
-def plot_causal_matrix(cmtx, class_names=None, figsize=None, vmin=None, vmax=None, show_text=True, cmap="magma"):
- """
- A function to create a colored and labeled causal matrix matplotlib figure
- given true labels and preds.
- Args:
- cmtx (ndarray): causal matrix.
- num_classes (int): total number of nodes.
- class_names (Optional[list of strs]): a list of node names.
- figsize (Optional[float, float]): the figure size of the causal matrix.
- If None, default to [6.4, 4.8].
-
- Returns:
- img (figure): matplotlib figure.
- """
- num_classes = cmtx.shape[0]
- if class_names is None or type(class_names) != list:
- class_names = [str(i) for i in range(num_classes)]
-
-
- figsize[0] = 30 if figsize[0] > 30 else figsize[0]
- figsize[1] = 20 if figsize[1] > 20 else figsize[1]
-
- plt.clf()
- plt.close("all")
- figure = plt.figure(figsize=figsize)
- plt.imshow(cmtx, interpolation="nearest",
- cmap=cmap, vmin=vmin, vmax=vmax)
- plt.title("Causal matrix")
- plt.colorbar()
- # tick_marks = np.arange(len(class_names))
- # plt.xticks(tick_marks, class_names, rotation=45)
- # plt.yticks(tick_marks, class_names)
-
- # Use white text if squares are dark; otherwise black.
- threshold = cmtx.max() / 2.0
- for i, j in itertools.product(range(cmtx.shape[0]), range(cmtx.shape[1])):
- color = "white" if cmtx[i, j] < threshold else "black"
- if cmtx.shape[0] < 20 and show_text:
- plt.text(j, i, format(cmtx[i, j], ".2e") if cmtx[i, j] != 0 else ".",
- horizontalalignment="center", color=color,)
-
- plt.tight_layout()
- plt.ylabel("True label")
- plt.xlabel("Predicted label")
-
- return figure
-
-
-def reproduc(seed, benchmark=False, deterministic=True):
- """Make experiments reproducible
- """
- random.seed(seed)
- np.random.seed(seed)
- torch.manual_seed(seed)
- torch.cuda.manual_seed_all(seed)
- torch.backends.cudnn.benchmark = benchmark
- torch.backends.cudnn.deterministic = deterministic
-
-
-def omegaconf2list(opt, prefix='', sep='.'):
- notation_list = []
- for k, v in opt.items():
- k = str(k)
- if isinstance(v, omegaconf.listconfig.ListConfig):
- notation_list.append("{}{}={}".format(prefix, k, v))
- # if k in ['iter_list','step_list']: # do not sparse list
- # dot_notation_list.append("{}{}={}".format(prefix,k,v))
- # else:
- # templist = []
- # for v_ in v:
- # templist.append('{}{}={}'.format(prefix,k,v_))
- # dot_notation_list.append(templist)
- elif isinstance(v, (float, str, int,)):
- notation_list.append("{}{}={}".format(prefix, k, v))
- elif v is None:
- notation_list.append("{}{}=~".format(prefix, k,))
- elif isinstance(v, omegaconf.dictconfig.DictConfig):
- nested_flat_list = omegaconf2list(v, prefix + k + sep, sep=sep)
- if nested_flat_list:
- notation_list.extend(nested_flat_list)
- else:
- raise NotImplementedError
- return notation_list
-
-
-def omegaconf2dotlist(opt, prefix='',):
- return omegaconf2list(opt, prefix, sep='.')
-
-
-def omegaconf2dict(opt, sep):
- notation_list = omegaconf2list(opt, sep=sep)
- dict = {notation.split('=', maxsplit=1)[0]: notation.split(
- '=', maxsplit=1)[1] for notation in notation_list}
- return dict
-
-
-# def read_video(video_path: str):
-# if ops(video_path)[-1] == ".tif":
-# data = tifffile.imread(video_path)
-# data = (data / np.max(data) * 255).astype(np.uint8)
-# if len(data.shape) == 3:
-# data = data[:, :, :, None]
-# return data
-# else:
-# cap = cv2.VideoCapture(video_path)
-# frames = []
-# while cap.isOpened():
-# # get a frame
-# ret, frame = cap.read()
-# if not ret:
-# break
-# frames.append(np.array(frame)[None])
-
-# cap.release()
-# return np.concatenate(frames, axis=0)
-
-
-# def save_video(video_path: str, data):
-# skvideo.io.vwrite(video_path, data)
-
-
-
-class LabelArray(object):
- # def __init__(self, array, labels):
- # self.arr = array
- # self.labels = labels
- # self.marks = dim_marks
- # assert [len[label_list] for label_list in labels] == self.arr.shape
-
- def __init__(self, dim, labels=None):
- if labels is not None:
- if len(dim) != dim:
- raise "The length of labels has to be equal to dim if defined"
- else:
- self.labels = deepcopy(labels)
- else:
- self.labels = [[] for _ in range(dim)]
- self.arr = None
- self.update_arr()
-
- def update_arr(self):
- if self.arr is not None:
- oldarr = self.arr
- self.arr = np.zeros([len(dim) for dim in self.labels]) * np.nan
- self.arr[tuple([slice(0,sh_dim,1) for sh_dim in oldarr.shape])] = oldarr
- else:
- self.arr = np.zeros([len(dim) for dim in self.labels]) * np.nan
- self.shape = self.arr.shape
-
-
- def __getitem__(self, label_list):
- index_list = []
- for dim,label in enumerate(label_list):
- if isinstance(label, str):
- index_list.append(self.labels[dim].index(label))
- elif isinstance(label, slice):
- index_list.append(label)
- elif isinstance(label, int):
- index_list.append(label)
- else:
- raise NotImplementedError
-
- return self.arr[tuple(index_list)]
-
- def __setitem__(self, label_list, val):
- index_list = []
- for dim,label in enumerate(label_list):
- if isinstance(label, str):
- if not label in self.labels[dim]:
- self.labels[dim].append(label)
- self.update_arr()
- index_list.append(self.labels[dim].index(label))
- elif isinstance(label, slice):
- index_list.append(label)
- elif isinstance(label, int):
- index_list.append(label)
- else:
- raise NotImplementedError
-
- self.arr[tuple(index_list)] = val
-
- def __str__(self):
- return str(self.arr) + "\n--------------------------\n" + str(self.labels)
-
-
- def to_np(self):
- return self.arr
-
- def from_np(self, array):
- assert [len[label_list] for label_list in self.labels] == self.arr.shape
- self.arr = array
-
diff --git a/singular_ticker_causal/algorithms/CUTS_PLUS/utils/opt_type.py b/singular_ticker_causal/algorithms/CUTS_PLUS/utils/opt_type.py
deleted file mode 100644
index b71880a7c614824e48c3f34301caeefbe599a159..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/algorithms/CUTS_PLUS/utils/opt_type.py
+++ /dev/null
@@ -1,83 +0,0 @@
-from dataclasses import dataclass
-from typing import Any
-
-@dataclass
-class ReproducOpt:
- seed: int
- benchmark: bool
- deterministic: bool
-
-@dataclass
-class NetworkOpt:
- name: str
- network_param: Any
-
-@dataclass
-class TrainOpt:
- batch_size: int
- total_epoch: int
- time_window: int
-
-
-@dataclass
-class TsGAEopt:
- dir_name: str
- task_name: str
- optimizer: Any
- reproduc: ReproducOpt
- network: NetworkOpt
- train: TrainOpt
- log: Any
- causal_thres: str
-
-@dataclass
-class MultiCADopt:
- dir_name: str
- task_name: str
-
- @dataclass
- class MultiCADargs:
- n_nodes: int
- input_step: int
- window_step: int
- stride: int
- batch_size: int
- sample_per_epoch: int
- data_dim: int
- total_epoch: int
-
- patience: int
- warmup: Any
-
- show_graph_every: int
- val_every: int
-
- n_groups: int
- group_policy: Any
- causal_thres: str
-
- @dataclass
- class data_pred:
- model: str
- merge_policy: str
- lr_data_start: float
- lr_data_end: float
- weight_decay: int
- prob: bool
-
- @dataclass
- class graph_discov:
- lr_graph_start: float
- lr_graph_end: float
- lambda_s_start: float
- lambda_s_end: float
- tau_start: float
- tau_end: float
- disable_bwd: bool
- separate_bwd: bool
- disable_ind: bool
- disable_graph: bool
- use_true_graph: bool
-
- reproduc: ReproducOpt
- log: Any
\ No newline at end of file
diff --git a/singular_ticker_causal/algorithms/__init__.py b/singular_ticker_causal/algorithms/__init__.py
deleted file mode 100644
index cd56e3269e7965566cc6e9ea7f195fd22b0df817..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/algorithms/__init__.py
+++ /dev/null
@@ -1,2 +0,0 @@
-# CAMEF algorithms subpackage — GPT4MTS and dataloader have been retired.
-# This package is intentionally empty pending removal of the CAMEF directory.
\ No newline at end of file
diff --git a/singular_ticker_causal/causal_inference/__init__.py b/singular_ticker_causal/causal_inference/__init__.py
deleted file mode 100644
index 315ca168a2d9ee302b992a46db9628053674dd11..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/causal_inference/__init__.py
+++ /dev/null
@@ -1,15 +0,0 @@
-from .causal_model import StructuralCausalModel
-from .identification import IdentificationStrategy, find_adjustment_set, is_identifiable
-from .pywhyllm_assumptions import CausalAssumptionReport, PyWhyLLMConfig, PyWhyLLMAssumptionService
-from .query_engine import CausalQueryEngine
-
-__all__ = [
- "CausalAssumptionReport",
- "CausalQueryEngine",
- "IdentificationStrategy",
- "PyWhyLLMConfig",
- "PyWhyLLMAssumptionService",
- "StructuralCausalModel",
- "find_adjustment_set",
- "is_identifiable",
-]
diff --git a/singular_ticker_causal/causal_inference/abduction.py b/singular_ticker_causal/causal_inference/abduction.py
deleted file mode 100644
index c21c2c06bbb81d19303b1cd09600d0a2d623504a..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/causal_inference/abduction.py
+++ /dev/null
@@ -1,127 +0,0 @@
-from __future__ import annotations
-
-from typing import Dict
-
-import numpy as np
-
-from .causal_model import StructuralCausalModel
-
-
-def abduct_noise(scm: StructuralCausalModel, observed_values: Dict[str, float]) -> Dict[str, float]:
- noise: Dict[str, float] = {}
- for node_idx in scm.topological_indices:
- node = scm.nodes[node_idx]
- if node not in observed_values:
- continue
- obs = float(observed_values[node])
- eq = scm.equations[node]
- if not eq.parents or eq.equation_type == "exogenous":
- noise[node] = obs - eq.intercept
- continue
- parent_vals = []
- for p in eq.parents:
- if p not in observed_values:
- parent_vals = []
- break
- parent_vals.append(float(observed_values[p]))
- if not parent_vals:
- continue
- coef_vec = np.array([eq.coefficients[p] for p in eq.parents], dtype=float)
- pred = eq.intercept + float(np.array(parent_vals, dtype=float) @ coef_vec)
- noise[node] = obs - pred
- return noise
-
-
-def counterfactual_predict(
- scm: StructuralCausalModel,
- observed_values: Dict[str, float],
- treatment: str,
- counterfactual_value: float,
- target: str,
-) -> Dict[str, float]:
- if treatment not in scm.node_to_idx:
- raise ValueError(f"Unknown treatment node: {treatment}")
- if target not in scm.node_to_idx:
- raise ValueError(f"Unknown target node: {target}")
-
- noises = abduct_noise(scm, observed_values)
- state = scm.data_level[-1].copy()
- for node, value in observed_values.items():
- if node in scm.node_to_idx:
- state[scm.node_to_idx[node]] = float(value)
-
- cf = state.copy()
- t_idx = scm.node_to_idx[treatment]
- cf[t_idx] = float(counterfactual_value)
-
- for node_idx in scm.topological_indices:
- node = scm.nodes[node_idx]
- if node_idx == t_idx:
- continue
- eq = scm.equations[node]
- if not eq.parents or eq.equation_type == "exogenous":
- cf[node_idx] = eq.intercept + noises.get(node, 0.0)
- continue
- parent_vals = cf[eq.parent_indices]
- coef_vec = np.array([eq.coefficients[p] for p in eq.parents], dtype=float)
- pred = eq.intercept + float(parent_vals @ coef_vec)
- cf[node_idx] = pred + noises.get(node, 0.0)
-
- y_idx = scm.node_to_idx[target]
- factual = float(state[y_idx])
- counterfactual = float(cf[y_idx])
- ite = counterfactual - factual
- pct_change = ite / (abs(factual) + 1e-12)
- return {
- "factual_outcome": factual,
- "counterfactual_outcome": counterfactual,
- "ite": float(ite),
- "pct_change": float(pct_change),
- }
-
-
-def counterfactual_predict_multi(
- scm: StructuralCausalModel,
- observed_values: Dict[str, float],
- counterfactual_values: Dict[str, float],
- target: str,
-) -> Dict[str, float]:
- """Compute a joint counterfactual for multiple simultaneous interventions."""
- for node in counterfactual_values:
- if node not in scm.node_to_idx:
- raise ValueError(f"Unknown counterfactual node: {node}")
-
- noises = abduct_noise(scm, observed_values)
- state = scm.data_level[-1].copy()
- for node, value in observed_values.items():
- if node in scm.node_to_idx:
- state[scm.node_to_idx[node]] = float(value)
-
- cf = state.copy()
- for node, value in counterfactual_values.items():
- cf[scm.node_to_idx[node]] = float(value)
-
- for node_idx in scm.topological_indices:
- node = scm.nodes[node_idx]
- if node in counterfactual_values:
- continue
- eq = scm.equations[node]
- if not eq.parents or eq.equation_type == "exogenous":
- cf[node_idx] = eq.intercept + noises.get(node, 0.0)
- continue
- parent_vals = cf[eq.parent_indices]
- coef_vec = np.array([eq.coefficients[p] for p in eq.parents], dtype=float)
- pred = eq.intercept + float(parent_vals @ coef_vec)
- cf[node_idx] = pred + noises.get(node, 0.0)
-
- y_idx = scm.node_to_idx[target]
- factual = float(state[y_idx])
- counterfactual = float(cf[y_idx])
- ite = counterfactual - factual
- pct_change = ite / (abs(factual) + 1e-12)
- return {
- "factual_outcome": factual,
- "counterfactual_outcome": counterfactual,
- "ite": float(ite),
- "pct_change": float(pct_change),
- }
diff --git a/singular_ticker_causal/causal_inference/causal_model.py b/singular_ticker_causal/causal_inference/causal_model.py
deleted file mode 100644
index bc7b5721f09fc2f97d58527b45217c6bfd5b1a07..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/causal_inference/causal_model.py
+++ /dev/null
@@ -1,314 +0,0 @@
-from __future__ import annotations
-
-from dataclasses import dataclass, field
-from typing import Dict, List, Optional, Sequence, Tuple
-
-import numpy as np
-
-
-@dataclass
-class StructuralEquation:
- node: str
- parents: List[str]
- parent_indices: List[int]
- intercept: float
- coefficients: Dict[str, float]
- residual_mean: float
- residual_std: float
- r_squared: float
- n_obs: int
- equation_type: str # "linear" | "exogenous"
-
-
-@dataclass
-class StructuralCausalModel:
- nodes: Sequence[str]
- adj: np.ndarray
- adjacency_mask: np.ndarray
- data_tech: np.ndarray
- mask_tech: Optional[np.ndarray] = None
- prohibition_mask: Optional[np.ndarray] = None
- threshold: float = 0.5
- lag: int = 1
- min_obs: int = 12
- ridge_alpha: float = 1e-4
- dag_adj: np.ndarray = field(init=False)
- topological_indices: List[int] = field(init=False, default_factory=list)
- equations: Dict[str, StructuralEquation] = field(init=False, default_factory=dict)
- removed_cycle_edges: List[Dict[str, float]] = field(init=False, default_factory=list)
-
- def __post_init__(self) -> None:
- self.nodes = list(self.nodes)
- self.node_to_idx = {node: i for i, node in enumerate(self.nodes)}
- self.data_level = self._extract_level_data(self.data_tech)
- self.mask_level = self._extract_mask_level(self.mask_tech)
- self._validate_shapes()
-
- @property
- def n_nodes(self) -> int:
- return len(self.nodes)
-
- @property
- def t_steps(self) -> int:
- return self.data_level.shape[0]
-
- def fit(self) -> "StructuralCausalModel":
- self.dag_adj = self._build_dag()
- self._enforce_acyclic()
- self.topological_indices = self._topological_sort(self.dag_adj)
- self._fit_equations()
- return self
-
- def density(self) -> float:
- n = self.n_nodes
- max_edges = n * (n - 1)
- if max_edges == 0:
- return 0.0
- return float(np.sum(self.dag_adj) / max_edges)
-
- def parents_of(self, node: str) -> List[str]:
- j = self.node_to_idx[node]
- return [self.nodes[i] for i in np.where(self.dag_adj[:, j])[0]]
-
- def has_directed_path(self, source: str, target: str) -> bool:
- s = self.node_to_idx[source]
- t = self.node_to_idx[target]
- stack = [s]
- visited = set()
- while stack:
- u = stack.pop()
- if u == t:
- return True
- if u in visited:
- continue
- visited.add(u)
- children = np.where(self.dag_adj[u])[0].tolist()
- stack.extend(children)
- return False
-
- def descendants_of(self, node: str) -> List[str]:
- start = self.node_to_idx[node]
- stack = [start]
- visited = set()
- while stack:
- u = stack.pop()
- children = np.where(self.dag_adj[u])[0].tolist()
- for v in children:
- if v not in visited:
- visited.add(v)
- stack.append(v)
- visited.discard(start)
- return [self.nodes[i] for i in sorted(visited)]
-
- def _validate_shapes(self) -> None:
- n = len(self.nodes)
- if self.adj.shape != (n, n):
- raise ValueError(f"adj shape mismatch: expected {(n, n)}, got {self.adj.shape}")
- if self.adjacency_mask.shape != (n, n):
- raise ValueError(
- f"adjacency_mask shape mismatch: expected {(n, n)}, got {self.adjacency_mask.shape}"
- )
- if self.prohibition_mask is not None and self.prohibition_mask.shape != (n, n):
- raise ValueError(
- f"prohibition_mask shape mismatch: expected {(n, n)}, got {self.prohibition_mask.shape}"
- )
- if self.data_level.ndim != 2 or self.data_level.shape[1] != n:
- raise ValueError(
- f"data_level shape mismatch: expected (T, {n}), got {self.data_level.shape}"
- )
- if self.mask_level is not None and self.mask_level.shape != self.data_level.shape:
- raise ValueError(
- "mask_tech shape mismatch after extraction: expected shape "
- f"{self.data_level.shape}, got {self.mask_level.shape}"
- )
- if self.lag < 1:
- raise ValueError("lag must be >= 1")
-
- def _extract_level_data(self, data_tech: np.ndarray) -> np.ndarray:
- if data_tech.ndim == 3:
- return np.asarray(data_tech[:, :, 0], dtype=float)
- if data_tech.ndim == 2:
- return np.asarray(data_tech, dtype=float)
- raise ValueError(f"Unsupported data_tech ndim={data_tech.ndim}; expected 2 or 3.")
-
- def _extract_mask_level(self, mask_tech: Optional[np.ndarray]) -> Optional[np.ndarray]:
- if mask_tech is None:
- return None
- if mask_tech.ndim == 3:
- return np.asarray(mask_tech[:, :, 0], dtype=float)
- if mask_tech.ndim == 2:
- return np.asarray(mask_tech, dtype=float)
- raise ValueError(f"Unsupported mask_tech ndim={mask_tech.ndim}; expected 2 or 3.")
-
- def _build_dag(self) -> np.ndarray:
- cuts_edges = self.adj >= self.threshold
- prior_edges = self.adjacency_mask > 0
- dag = np.logical_or(cuts_edges, prior_edges)
- if self.prohibition_mask is not None:
- prohibited = self.prohibition_mask <= 0
- dag = np.where(prohibited, False, dag)
- np.fill_diagonal(dag, False)
- return dag.astype(bool)
-
- def _enforce_acyclic(self) -> None:
- while True:
- cycle_edges = self._find_cycle_edges(self.dag_adj)
- if not cycle_edges:
- return
-
- removable = []
- for src, dst in cycle_edges:
- mandatory = bool(self.adjacency_mask[src, dst] > 0)
- score = float(self.adj[src, dst])
- removable.append((mandatory, score, src, dst))
-
- non_mandatory = [r for r in removable if not r[0]]
- choice = min(non_mandatory or removable, key=lambda x: x[1])
- _, score, src, dst = choice
-
- self.dag_adj[src, dst] = False
- self.removed_cycle_edges.append(
- {
- "source": self.nodes[src],
- "target": self.nodes[dst],
- "adj_score": score,
- }
- )
-
- def _topological_sort(self, dag_adj: np.ndarray) -> List[int]:
- n = dag_adj.shape[0]
- indegree = np.sum(dag_adj, axis=0).astype(int)
- queue = [i for i in range(n) if indegree[i] == 0]
- order: List[int] = []
-
- while queue:
- node = queue.pop(0)
- order.append(node)
- children = np.where(dag_adj[node])[0]
- for child in children:
- indegree[child] -= 1
- if indegree[child] == 0:
- queue.append(int(child))
-
- if len(order) != n:
- raise RuntimeError("DAG still contains a cycle after pruning.")
- return order
-
- def _find_cycle_edges(self, dag_adj: np.ndarray) -> List[Tuple[int, int]]:
- n = dag_adj.shape[0]
- state = np.zeros(n, dtype=int) # 0=unvisited, 1=visiting, 2=done
- parent = -np.ones(n, dtype=int)
-
- def dfs(u: int) -> Optional[List[Tuple[int, int]]]:
- state[u] = 1
- for v in np.where(dag_adj[u])[0]:
- v = int(v)
- if state[v] == 0:
- parent[v] = u
- found = dfs(v)
- if found:
- return found
- elif state[v] == 1:
- nodes = [v]
- cur = u
- while cur != v and cur != -1:
- nodes.append(cur)
- cur = int(parent[cur])
- nodes.append(v)
- nodes.reverse()
- return [(nodes[i], nodes[i + 1]) for i in range(len(nodes) - 1)]
- state[u] = 2
- return None
-
- for start in range(n):
- if state[start] == 0:
- result = dfs(start)
- if result:
- return result
- return []
-
- def _fit_equations(self) -> None:
- T = self.t_steps
- for idx in self.topological_indices:
- node = self.nodes[idx]
- parent_idx = [int(i) for i in np.where(self.dag_adj[:, idx])[0]]
- parent_names = [self.nodes[i] for i in parent_idx]
- valid_t = self._valid_timesteps(idx, parent_idx)
-
- if len(valid_t) == 0:
- self.equations[node] = StructuralEquation(
- node=node,
- parents=parent_names,
- parent_indices=parent_idx,
- intercept=0.0,
- coefficients={},
- residual_mean=0.0,
- residual_std=1.0,
- r_squared=0.0,
- n_obs=0,
- equation_type="exogenous",
- )
- continue
-
- y = self.data_level[valid_t, idx]
-
- if not parent_idx or len(valid_t) < self.min_obs:
- mu = float(np.mean(y))
- residuals = y - mu
- self.equations[node] = StructuralEquation(
- node=node,
- parents=parent_names,
- parent_indices=parent_idx,
- intercept=mu,
- coefficients={},
- residual_mean=float(np.mean(residuals)) if residuals.size else 0.0,
- residual_std=float(np.std(residuals)) if residuals.size else 1.0,
- r_squared=0.0,
- n_obs=int(len(valid_t)),
- equation_type="exogenous",
- )
- continue
-
- X = self.data_level[valid_t - self.lag][:, parent_idx]
- coef, intercept = self._fit_ridge(X, y)
- y_hat = intercept + (X @ coef)
- residuals = y - y_hat
- ss_res = float(np.sum(residuals ** 2))
- ss_tot = float(np.sum((y - np.mean(y)) ** 2))
- r2 = 1.0 - (ss_res / ss_tot) if ss_tot > 1e-12 else 0.0
-
- coeffs = {name: float(coef[i]) for i, name in enumerate(parent_names)}
- self.equations[node] = StructuralEquation(
- node=node,
- parents=parent_names,
- parent_indices=parent_idx,
- intercept=float(intercept),
- coefficients=coeffs,
- residual_mean=float(np.mean(residuals)),
- residual_std=float(np.std(residuals)),
- r_squared=r2,
- n_obs=int(len(valid_t)),
- equation_type="linear",
- )
-
- def _valid_timesteps(self, node_idx: int, parent_indices: List[int]) -> np.ndarray:
- valid_t = np.arange(self.lag, self.t_steps, dtype=int)
- if self.mask_level is None:
- return valid_t
-
- valid = self.mask_level[valid_t, node_idx] > 0
- for p_idx in parent_indices:
- valid &= self.mask_level[valid_t - self.lag, p_idx] > 0
- return valid_t[valid]
-
- def _fit_ridge(self, X: np.ndarray, y: np.ndarray) -> Tuple[np.ndarray, float]:
- X_mean = np.mean(X, axis=0)
- y_mean = float(np.mean(y))
- Xc = X - X_mean
- yc = y - y_mean
-
- p = X.shape[1]
- reg = self.ridge_alpha * np.eye(p)
- beta = np.linalg.solve(Xc.T @ Xc + reg, Xc.T @ yc)
- intercept = y_mean - float(X_mean @ beta)
- return beta, intercept
diff --git a/singular_ticker_causal/causal_inference/estimator.py b/singular_ticker_causal/causal_inference/estimator.py
deleted file mode 100644
index b8542108b0cb4abf1173421745290f996a58a139..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/causal_inference/estimator.py
+++ /dev/null
@@ -1,125 +0,0 @@
-from __future__ import annotations
-
-from typing import Dict, Optional, Set
-
-import numpy as np
-
-from .causal_model import StructuralCausalModel
-from .mutilator import propagate_intervention
-
-
-def estimate_ate(
- scm: StructuralCausalModel,
- treatment: str,
- outcome: str,
- treatment_value: Optional[float] = None,
- adjustment_set: Optional[Set[str]] = None,
- horizon: int = 5,
-) -> Dict[str, object]:
- t_idx = scm.node_to_idx[treatment]
- y_idx = scm.node_to_idx[outcome]
- z_nodes = sorted(adjustment_set or set())
- z_idx = [scm.node_to_idx[z] for z in z_nodes]
-
- valid_t = np.arange(scm.lag, scm.t_steps, dtype=int)
- if scm.mask_level is not None:
- valid = (scm.mask_level[valid_t, y_idx] > 0) & (scm.mask_level[valid_t - scm.lag, t_idx] > 0)
- for zi in z_idx:
- valid &= scm.mask_level[valid_t - scm.lag, zi] > 0
- valid_t = valid_t[valid]
-
- if len(valid_t) < 5:
- raise ValueError(
- f"Insufficient observations for ATE estimation of {treatment}->{outcome}: {len(valid_t)} rows."
- )
-
- y = scm.data_level[valid_t, y_idx]
- x_treat = scm.data_level[valid_t - scm.lag, t_idx]
- X_parts = [np.ones((len(valid_t), 1)), x_treat.reshape(-1, 1)]
- if z_idx:
- X_parts.append(scm.data_level[valid_t - scm.lag][:, z_idx])
- X = np.concatenate(X_parts, axis=1)
-
- beta = np.linalg.pinv(X.T @ X) @ (X.T @ y)
- y_hat = X @ beta
- resid = y - y_hat
- dof = max(1, len(y) - X.shape[1])
- sigma2 = float(np.sum(resid ** 2) / dof)
- cov = sigma2 * np.linalg.pinv(X.T @ X)
- se = float(np.sqrt(max(cov[1, 1], 0.0)))
- ate = float(beta[1])
-
- baseline = float(np.mean(np.abs(y))) + 1e-12
- ate_normalized = ate / baseline
- ci_95 = (ate - 1.96 * se, ate + 1.96 * se)
-
- path_contributions = _path_contributions(scm, treatment, outcome)
-
- horizon_effects = []
- if treatment_value is not None:
- baseline_t = scm.t_steps - 1
- shocked = propagate_intervention(scm, treatment, treatment_value, horizon=horizon, baseline_t=baseline_t)
- base_outcome = float(scm.data_level[baseline_t, y_idx])
- for h in range(horizon):
- horizon_effects.append(float(shocked[h, y_idx] - base_outcome))
-
- return {
- "ate": ate,
- "ate_normalized": float(ate_normalized),
- "ci_95": (float(ci_95[0]), float(ci_95[1])),
- "n_obs": int(len(valid_t)),
- "path_contributions": path_contributions,
- "horizon_effects": horizon_effects,
- }
-
-
-def _path_contributions(
- scm: StructuralCausalModel,
- treatment: str,
- outcome: str,
- max_paths: int = 100,
-) -> Dict[str, float]:
- start = scm.node_to_idx[treatment]
- target = scm.node_to_idx[outcome]
- paths = []
-
- def dfs(node: int, path: list[int], seen: set[int]) -> None:
- if len(paths) >= max_paths:
- return
- if node == target:
- paths.append(path.copy())
- return
- for child in np.where(scm.dag_adj[node])[0]:
- child = int(child)
- if child in seen:
- continue
- seen.add(child)
- path.append(child)
- dfs(child, path, seen)
- path.pop()
- seen.remove(child)
-
- dfs(start, [start], {start})
-
- contributions: Dict[str, float] = {}
- for path in paths:
- coeff_product = 1.0
- valid = True
- for i in range(len(path) - 1):
- src = scm.nodes[path[i]]
- dst = scm.nodes[path[i + 1]]
- eq = scm.equations.get(dst)
- if eq is None:
- valid = False
- break
- coef = eq.coefficients.get(src)
- if coef is None:
- valid = False
- break
- coeff_product *= coef
- if not valid:
- continue
- label = " -> ".join(scm.nodes[i] for i in path)
- contributions[label] = float(coeff_product)
-
- return contributions
diff --git a/singular_ticker_causal/causal_inference/identification.py b/singular_ticker_causal/causal_inference/identification.py
deleted file mode 100644
index 0752702f4e9472763749a776eda865c81927d745..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/causal_inference/identification.py
+++ /dev/null
@@ -1,41 +0,0 @@
-from __future__ import annotations
-
-from enum import Enum
-from typing import Optional, Set, Tuple
-
-from .causal_model import StructuralCausalModel
-
-
-class IdentificationStrategy(Enum):
- DIRECT = "direct"
- BACKDOOR = "backdoor"
- NOT_IDENTIFIABLE = "not_identifiable"
-
-
-def find_adjustment_set(
- scm: StructuralCausalModel,
- treatment: str,
- outcome: str,
-) -> Tuple[IdentificationStrategy, Optional[Set[str]]]:
- if treatment not in scm.node_to_idx:
- raise ValueError(f"Unknown treatment node: {treatment}")
- if outcome not in scm.node_to_idx:
- raise ValueError(f"Unknown outcome node: {outcome}")
- if treatment == outcome:
- return IdentificationStrategy.DIRECT, set()
-
- if not scm.has_directed_path(treatment, outcome):
- return IdentificationStrategy.NOT_IDENTIFIABLE, None
-
- parents = set(scm.parents_of(treatment))
- descendants = set(scm.descendants_of(treatment))
- adjustment = parents - descendants - {outcome}
-
- if adjustment:
- return IdentificationStrategy.BACKDOOR, adjustment
- return IdentificationStrategy.DIRECT, set()
-
-
-def is_identifiable(scm: StructuralCausalModel, treatment: str, outcome: str) -> bool:
- strategy, _ = find_adjustment_set(scm, treatment, outcome)
- return strategy != IdentificationStrategy.NOT_IDENTIFIABLE
diff --git a/singular_ticker_causal/causal_inference/mutilator.py b/singular_ticker_causal/causal_inference/mutilator.py
deleted file mode 100644
index 49330d5014ed24cc8ffe86ba523b87b7275fccf6..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/causal_inference/mutilator.py
+++ /dev/null
@@ -1,62 +0,0 @@
-from __future__ import annotations
-
-from copy import deepcopy
-from typing import Dict, Optional, Sequence
-
-import numpy as np
-
-from .causal_model import StructuralCausalModel
-
-
-def mutilate_graph(scm: StructuralCausalModel, treatment: str, value: float) -> StructuralCausalModel:
- mutilated = deepcopy(scm)
- t_idx = mutilated.node_to_idx[treatment]
- mutilated.dag_adj[:, t_idx] = False
- mutilated.pinned_values = {treatment: float(value)}
- return mutilated
-
-
-def propagate_intervention(
- scm: StructuralCausalModel,
- treatment: str,
- value: float,
- targets: Optional[Sequence[str]] = None,
- horizon: int = 5,
- baseline_t: int = -1,
-) -> np.ndarray:
- if horizon < 1:
- raise ValueError("horizon must be >= 1")
-
- baseline_idx = baseline_t if baseline_t >= 0 else (scm.t_steps + baseline_t)
- if baseline_idx < 0 or baseline_idx >= scm.t_steps:
- raise ValueError(f"baseline_t {baseline_t} resolves out of bounds for T={scm.t_steps}")
-
- t_idx = scm.node_to_idx[treatment]
- prev = scm.data_level[baseline_idx].copy()
- forecasts = np.zeros((horizon, scm.n_nodes), dtype=float)
-
- for h in range(horizon):
- nxt = prev.copy()
- nxt[t_idx] = float(value)
- for node_idx in scm.topological_indices:
- if node_idx == t_idx:
- continue
- node = scm.nodes[node_idx]
- eq = scm.equations[node]
- if not eq.parent_indices or eq.equation_type == "exogenous":
- continue
- parent_vals = prev[eq.parent_indices]
- coef_vec = np.array([eq.coefficients[p] for p in eq.parents], dtype=float)
- nxt[node_idx] = eq.intercept + float(parent_vals @ coef_vec)
- forecasts[h] = nxt
- prev = nxt
-
- if targets:
- missing = [n for n in targets if n not in scm.node_to_idx]
- if missing:
- raise ValueError(f"Unknown targets: {missing}")
- return forecasts
-
-
-def target_series(forecasts: np.ndarray, scm: StructuralCausalModel, targets: Sequence[str]) -> Dict[str, list[float]]:
- return {t: [float(v) for v in forecasts[:, scm.node_to_idx[t]]] for t in targets}
diff --git a/singular_ticker_causal/causal_inference/pywhyllm_assumptions.py b/singular_ticker_causal/causal_inference/pywhyllm_assumptions.py
deleted file mode 100644
index 291d2f45dbf05e2aa27f859ef43e69bc13d720a6..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/causal_inference/pywhyllm_assumptions.py
+++ /dev/null
@@ -1,338 +0,0 @@
-from __future__ import annotations
-
-import hashlib
-import json
-import os
-from dataclasses import asdict, dataclass, field
-from typing import Any, Dict, Iterable, List, Optional, Sequence, Tuple
-
-import numpy as np
-
-
-def _json_safe(value: Any) -> Any:
- if isinstance(value, np.ndarray):
- return value.tolist()
- if isinstance(value, (np.integer,)):
- return int(value)
- if isinstance(value, (np.floating,)):
- return float(value)
- if isinstance(value, dict):
- return {str(k): _json_safe(v) for k, v in value.items()}
- if isinstance(value, (list, tuple, set)):
- return [_json_safe(v) for v in value]
- if isinstance(value, (str, int, float, bool)) or value is None:
- return value
- return str(value)
-
-
-@dataclass
-class PyWhyLLMConfig:
- enabled: bool = False
- model: str = "gpt-4"
- max_edges: int = 25
- cache_dir: str = os.path.join(
- os.path.dirname(__file__),
- "..",
- "debug_data",
- "pywhyllm_cache",
- )
-
-
-@dataclass
-class CausalAssumptionReport:
- available: bool = True
- reason: Optional[str] = None
- domain_expertises: List[str] = field(default_factory=list)
- suggested_confounders: List[str] = field(default_factory=list)
- suggested_backdoor_sets: List[List[str]] = field(default_factory=list)
- suggested_mediators: List[str] = field(default_factory=list)
- suggested_ivs: List[str] = field(default_factory=list)
- negative_controls: List[str] = field(default_factory=list)
- latent_confounders: List[str] = field(default_factory=list)
- edge_critiques: Any = field(default_factory=list)
- accepted_edges: List[Tuple[str, str]] = field(default_factory=list)
- rejected_edges: List[Tuple[str, str]] = field(default_factory=list)
- warnings: List[str] = field(default_factory=list)
-
- def to_dict(self) -> Dict[str, Any]:
- payload = asdict(self)
- payload["accepted_edges"] = [list(edge) for edge in self.accepted_edges]
- payload["rejected_edges"] = [list(edge) for edge in self.rejected_edges]
- return _json_safe(payload)
-
- @classmethod
- def from_dict(cls, payload: Dict[str, Any]) -> "CausalAssumptionReport":
- data = dict(payload)
- data["accepted_edges"] = [tuple(edge) for edge in data.get("accepted_edges", [])]
- data["rejected_edges"] = [tuple(edge) for edge in data.get("rejected_edges", [])]
- return cls(**data)
-
-
-def _dedupe_strings(values: Iterable[Any]) -> List[str]:
- seen = set()
- result: List[str] = []
- for value in values or []:
- text = str(value).strip()
- if text and text not in seen:
- seen.add(text)
- result.append(text)
- return result
-
-
-def _normalise_suggestion(value: Any) -> List[str]:
- if value is None:
- return []
- if isinstance(value, tuple) and len(value) == 2:
- return _normalise_suggestion(value[1])
- if isinstance(value, dict):
- keys = [k for k, v in value.items() if isinstance(k, str) and v]
- if keys:
- return _dedupe_strings(keys)
- flattened: List[str] = []
- for item in value.values():
- flattened.extend(_normalise_suggestion(item))
- return _dedupe_strings(flattened)
- if isinstance(value, (list, set, tuple)):
- flattened = []
- for item in value:
- if isinstance(item, (list, set, tuple, dict)):
- flattened.extend(_normalise_suggestion(item))
- else:
- flattened.append(item)
- return _dedupe_strings(flattened)
- return _dedupe_strings([value])
-
-
-class PyWhyLLMAssumptionService:
- def __init__(
- self,
- config: Optional[PyWhyLLMConfig] = None,
- model_suggester: Any = None,
- identification_suggester: Any = None,
- validation_suggester: Any = None,
- relationship_strategy: Any = None,
- ):
- self.config = config or PyWhyLLMConfig()
- self._model_suggester = model_suggester
- self._identification_suggester = identification_suggester
- self._validation_suggester = validation_suggester
- self._relationship_strategy = relationship_strategy
-
- @property
- def enabled(self) -> bool:
- return bool(self.config.enabled)
-
- def analyze(
- self,
- *,
- nodes: Sequence[str],
- dag_adj: np.ndarray,
- treatment: str,
- outcome: str,
- max_edges: Optional[int] = None,
- ) -> CausalAssumptionReport:
- if not self.enabled:
- return CausalAssumptionReport(
- available=False,
- reason="PyWhy-LLM is disabled. Set PYWHYLLM_ENABLED=true or pass pywhyllm_enabled=True.",
- )
-
- cache_path = self._cache_path(nodes, dag_adj, treatment, outcome)
- cached = self._read_cache(cache_path)
- if cached is not None:
- return cached
-
- try:
- modeler, identifier, validator, relationship_strategy = self._suggesters()
- except Exception as exc:
- return CausalAssumptionReport(
- available=False,
- reason=f"PyWhy-LLM is not installed or failed to initialize: {exc}",
- )
-
- all_factors = list(nodes)
- edges = self._edges(nodes, dag_adj, max_edges or self.config.max_edges)
- warnings: List[str] = []
-
- domain_expertises: List[str] = []
- suggested_confounders: List[str] = []
- suggested_backdoor_sets: List[List[str]] = []
- suggested_mediators: List[str] = []
- suggested_ivs: List[str] = []
- negative_controls: List[str] = []
- latent_confounders: List[str] = []
- edge_critiques: Any = []
- suggested_dag: Any = edges
-
- try:
- domain_expertises = _normalise_suggestion(modeler.suggest_domain_expertises(all_factors))
- except Exception as exc:
- warnings.append(f"domain_expertises failed: {exc}")
-
- try:
- suggested_confounders = _normalise_suggestion(
- modeler.suggest_confounders(treatment, outcome, all_factors, domain_expertises)
- )
- except Exception as exc:
- warnings.append(f"confounder suggestion failed: {exc}")
-
- try:
- suggested_dag = modeler.suggest_relationships(
- treatment,
- outcome,
- all_factors,
- domain_expertises,
- relationship_strategy,
- )
- except Exception as exc:
- warnings.append(f"relationship suggestion failed: {exc}")
-
- try:
- backdoor = identifier.suggest_backdoor(treatment, outcome, all_factors, domain_expertises)
- backdoor_nodes = _normalise_suggestion(backdoor)
- if backdoor_nodes:
- suggested_backdoor_sets = [backdoor_nodes]
- except Exception as exc:
- warnings.append(f"backdoor suggestion failed: {exc}")
-
- try:
- suggested_mediators = _normalise_suggestion(
- identifier.suggest_mediators(treatment, outcome, all_factors, domain_expertises)
- )
- except Exception as exc:
- warnings.append(f"mediator suggestion failed: {exc}")
-
- try:
- suggested_ivs = _normalise_suggestion(
- identifier.suggest_ivs(treatment, outcome, all_factors, domain_expertises)
- )
- except Exception as exc:
- warnings.append(f"iv suggestion failed: {exc}")
-
- try:
- edge_critiques = validator.critique_graph(
- all_factors,
- suggested_dag,
- domain_expertises,
- relationship_strategy,
- )
- except Exception as exc:
- warnings.append(f"edge critique failed: {exc}")
-
- try:
- latent_confounders = _normalise_suggestion(
- validator.suggest_latent_confounders(treatment, outcome, all_factors, domain_expertises)
- )
- except Exception as exc:
- warnings.append(f"latent confounder suggestion failed: {exc}")
-
- try:
- negative_controls = _normalise_suggestion(
- validator.suggest_negative_controls(treatment, outcome, all_factors, domain_expertises)
- )
- except Exception as exc:
- warnings.append(f"negative control suggestion failed: {exc}")
-
- accepted_edges, rejected_edges = self._classify_edges(edges, edge_critiques)
- report = CausalAssumptionReport(
- available=True,
- domain_expertises=domain_expertises,
- suggested_confounders=suggested_confounders,
- suggested_backdoor_sets=suggested_backdoor_sets,
- suggested_mediators=suggested_mediators,
- suggested_ivs=suggested_ivs,
- negative_controls=negative_controls,
- latent_confounders=latent_confounders,
- edge_critiques=edge_critiques,
- accepted_edges=accepted_edges,
- rejected_edges=rejected_edges,
- warnings=warnings,
- )
- self._write_cache(cache_path, report)
- return report
-
- def _suggesters(self) -> Tuple[Any, Any, Any, Any]:
- if self._model_suggester and self._identification_suggester and self._validation_suggester:
- return (
- self._model_suggester,
- self._identification_suggester,
- self._validation_suggester,
- self._relationship_strategy,
- )
-
- from pywhyllm import RelationshipStrategy
- from pywhyllm.suggesters.identification_suggester import IdentificationSuggester
- from pywhyllm.suggesters.model_suggester import ModelSuggester
- from pywhyllm.suggesters.validation_suggester import ValidationSuggester
-
- relationship_strategy = self._relationship_strategy or RelationshipStrategy.Pairwise
- return (
- self._model_suggester or ModelSuggester(self.config.model),
- self._identification_suggester or IdentificationSuggester(self.config.model),
- self._validation_suggester or ValidationSuggester(self.config.model),
- relationship_strategy,
- )
-
- def _cache_path(
- self,
- nodes: Sequence[str],
- dag_adj: np.ndarray,
- treatment: str,
- outcome: str,
- ) -> str:
- payload = {
- "nodes": list(nodes),
- "dag_adj": np.asarray(dag_adj, dtype=int).tolist(),
- "treatment": treatment,
- "outcome": outcome,
- "model": self.config.model,
- }
- digest = hashlib.sha256(json.dumps(payload, sort_keys=True).encode("utf-8")).hexdigest()
- return os.path.join(os.path.abspath(self.config.cache_dir), f"{digest}.json")
-
- def _read_cache(self, path: str) -> Optional[CausalAssumptionReport]:
- if not os.path.exists(path):
- return None
- try:
- with open(path) as f:
- return CausalAssumptionReport.from_dict(json.load(f))
- except Exception:
- return None
-
- def _write_cache(self, path: str, report: CausalAssumptionReport) -> None:
- try:
- os.makedirs(os.path.dirname(path), exist_ok=True)
- with open(path, "w") as f:
- json.dump(report.to_dict(), f, indent=2)
- except Exception:
- pass
-
- def _edges(self, nodes: Sequence[str], dag_adj: np.ndarray, max_edges: int) -> List[Tuple[str, str]]:
- found: List[Tuple[str, str, float]] = []
- for src_idx, src in enumerate(nodes):
- for dst_idx, dst in enumerate(nodes):
- if bool(dag_adj[src_idx, dst_idx]):
- found.append((src, dst, float(dag_adj[src_idx, dst_idx])))
- found.sort(key=lambda edge: abs(edge[2]), reverse=True)
- return [(src, dst) for src, dst, _ in found[:max_edges]]
-
- def _classify_edges(
- self,
- edges: List[Tuple[str, str]],
- edge_critiques: Any,
- ) -> Tuple[List[Tuple[str, str]], List[Tuple[str, str]]]:
- critique_text = str(edge_critiques).lower()
- rejected: List[Tuple[str, str]] = []
- for edge in edges:
- src, dst = edge
- edge_tokens = [
- f"{src}->{dst}".lower(),
- f"{src} -> {dst}".lower(),
- f"{src}, {dst}".lower(),
- ]
- if any(token in critique_text for token in edge_tokens) and any(
- marker in critique_text for marker in ["reject", "unlikely", "invalid", "implausible"]
- ):
- rejected.append(edge)
- accepted = [edge for edge in edges if edge not in rejected]
- return accepted, rejected
diff --git a/singular_ticker_causal/causal_inference/query_engine.py b/singular_ticker_causal/causal_inference/query_engine.py
deleted file mode 100644
index b5196c5917b49eeb83e401cdbdaad3bbfba8a422..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/causal_inference/query_engine.py
+++ /dev/null
@@ -1,505 +0,0 @@
-from __future__ import annotations
-
-import os
-import re
-from typing import Any, Dict, List, Optional, Sequence
-
-import numpy as np
-import pandas as pd
-
-from .abduction import counterfactual_predict, counterfactual_predict_multi
-from .causal_model import StructuralCausalModel
-from .estimator import estimate_ate
-from .identification import find_adjustment_set
-from .mutilator import mutilate_graph, propagate_intervention, target_series
-from .pywhyllm_assumptions import PyWhyLLMConfig, PyWhyLLMAssumptionService
-
-
-class CausalQueryEngine:
- def __init__(
- self,
- scm: StructuralCausalModel,
- data_tech: Optional[np.ndarray] = None,
- data_text: Optional[np.ndarray] = None,
- pywhyllm_service: Optional[PyWhyLLMAssumptionService] = None,
- pywhyllm_enabled: bool = False,
- ):
- self.scm = scm
- self.data_tech = data_tech
- self.data_text = data_text
- self.pywhyllm_enabled = pywhyllm_enabled
- self.pywhyllm_service = pywhyllm_service
-
- def assert_edge(self, treatment: str, outcome: str) -> Dict[str, object]:
- strategy, z = find_adjustment_set(self.scm, treatment, outcome)
- if strategy.value == "not_identifiable":
- return {
- "ate": 0.0,
- "ci_95": (0.0, 0.0),
- "strategy": strategy.name,
- "adjustment_set": [],
- "identifiable": False,
- "path_contributions": {},
- }
-
- est = estimate_ate(
- self.scm,
- treatment=treatment,
- outcome=outcome,
- adjustment_set=z,
- treatment_value=None,
- )
- return {
- "ate": est["ate"],
- "ci_95": est["ci_95"],
- "strategy": strategy.name,
- "adjustment_set": sorted(z or set()),
- "identifiable": True,
- "path_contributions": est["path_contributions"],
- }
-
- def intervene(self, treatment: str, value: float, targets: List[str], horizon: int = 5) -> Dict[str, object]:
- mutilated = mutilate_graph(self.scm, treatment=treatment, value=value)
- forecasts = propagate_intervention(
- mutilated,
- treatment=treatment,
- value=value,
- targets=targets,
- horizon=horizon,
- baseline_t=-1,
- )
- predicted = target_series(forecasts, self.scm, targets)
- ate_per_target = {}
- for target in targets:
- tidx = self.scm.node_to_idx[target]
- base = float(self.scm.data_level[-1, tidx])
- ate_per_target[target] = float(forecasts[0, tidx] - base)
- return {
- "mutilated_adj": mutilated.dag_adj.copy(),
- "predicted_values": predicted,
- "ate_per_target": ate_per_target,
- "horizon": horizon,
- }
-
- def _counterfactual_outcome(
- self,
- observed: Dict[str, float],
- counterfactual_values: Dict[str, float],
- target: str,
- ) -> Dict[str, float]:
- if len(counterfactual_values) == 1:
- treatment, value = next(iter(counterfactual_values.items()))
- return counterfactual_predict(
- self.scm,
- observed_values=observed,
- treatment=treatment,
- counterfactual_value=value,
- target=target,
- )
- return counterfactual_predict_multi(
- self.scm,
- observed_values=observed,
- counterfactual_values=counterfactual_values,
- target=target,
- )
-
- def _shapley_contributions(
- self,
- observed: Dict[str, float],
- interventions: Dict[str, float],
- target: str,
- mc_samples: int = 1000,
- n_jobs: int = -1,
- ) -> Dict[str, float]:
- import random
- from math import factorial
- from joblib import Parallel, delayed
- import os
-
- treatments = list(interventions.keys())
- n = len(treatments)
- if n == 1:
- return {treatments[0]: self._counterfactual_outcome(observed, interventions, target)["ite"]}
-
- contributions: Dict[str, float] = {t: 0.0 for t in treatments}
-
- # Use exact if n <= 10, else Monte Carlo
- use_exact = n <= 10
- if n_jobs < 0:
- n_jobs = os.cpu_count() or 4
-
- def marginal_contribution(k: str, subset: list[str]) -> float:
- with_k = {**{t: interventions[t] for t in subset}, k: interventions[k]}
- without_k = {t: interventions[t] for t in subset}
- v_with = self._counterfactual_outcome(observed, with_k, target)["counterfactual_outcome"]
- v_without = self._counterfactual_outcome(observed, without_k, target)["counterfactual_outcome"]
- return v_with - v_without
-
- if use_exact:
- all_factorial = float(factorial(n))
- tasks = []
-
- for k in treatments:
- others = [t for t in treatments if t != k]
- for r in range(len(others) + 1):
- for subset in __import__("itertools").combinations(others, r):
- subset_list = list(subset)
- weight = float(factorial(len(subset_list)) * factorial(n - len(subset_list) - 1) / all_factorial)
- tasks.append((k, subset_list, weight))
-
- results = Parallel(n_jobs=n_jobs, backend="threading")(
- delayed(marginal_contribution)(task[0], task[1]) for task in tasks
- )
- for task, res in zip(tasks, results):
- k, _, weight = task
- contributions[k] += weight * res
- else:
- # Monte Carlo approximation
- tasks = []
- for _ in range(mc_samples):
- perm = treatments.copy()
- random.shuffle(perm)
- for i, k in enumerate(perm):
- subset_list = perm[:i]
- tasks.append((k, subset_list))
-
- results = Parallel(n_jobs=n_jobs, backend="threading")(
- delayed(marginal_contribution)(task[0], task[1]) for task in tasks
- )
- for task, res in zip(tasks, results):
- k, _ = task
- contributions[k] += res / mc_samples
-
- return contributions
-
- def _build_dowhy_graph(self) -> str:
- edges = []
- for src_idx, src in enumerate(self.scm.nodes):
- for dst_idx, dst in enumerate(self.scm.nodes):
- if self.scm.dag_adj[src_idx, dst_idx]:
- edges.append(f"{src} -> {dst}")
- return "digraph{" + "; ".join(edges) + "}"
-
- def _default_pywhyllm_service(self) -> PyWhyLLMAssumptionService:
- enabled = self.pywhyllm_enabled or os.environ.get("PYWHYLLM_ENABLED", "").lower() in {
- "1",
- "true",
- "yes",
- "on",
- }
- return PyWhyLLMAssumptionService(
- PyWhyLLMConfig(
- enabled=enabled,
- model=os.environ.get("PYWHYLLM_MODEL", "gpt-4"),
- max_edges=int(os.environ.get("PYWHYLLM_MAX_EDGES", "25")),
- cache_dir=os.environ.get(
- "PYWHYLLM_CACHE_DIR",
- PyWhyLLMConfig().cache_dir,
- ),
- )
- )
-
- def _get_pywhyllm_service(self) -> PyWhyLLMAssumptionService:
- if self.pywhyllm_service is None:
- self.pywhyllm_service = self._default_pywhyllm_service()
- return self.pywhyllm_service
-
- def analyze_assumptions_with_pywhyllm(
- self,
- treatment: str,
- outcome: str,
- max_edges: Optional[int] = None,
- ) -> Dict[str, object]:
- report = self._get_pywhyllm_service().analyze(
- nodes=self.scm.nodes,
- dag_adj=self.scm.dag_adj,
- treatment=treatment,
- outcome=outcome,
- max_edges=max_edges,
- )
- return report.to_dict()
-
- def _valid_nodes(self, candidates: Sequence[Any]) -> List[str]:
- valid = set(self.scm.nodes)
- result = []
- for candidate in candidates or []:
- node = str(candidate).strip()
- if node in valid and node not in result:
- result.append(node)
- return result
-
- def _valid_backdoor_sets(self, report: Dict[str, Any], treatment: str, outcome: str) -> List[List[str]]:
- result = []
- blocked = {treatment, outcome}
- for suggested_set in report.get("suggested_backdoor_sets") or []:
- valid_set = [node for node in self._valid_nodes(suggested_set) if node not in blocked]
- if valid_set and valid_set not in result:
- result.append(valid_set)
- confounders = [node for node in self._valid_nodes(report.get("suggested_confounders") or []) if node not in blocked]
- if confounders and confounders not in result:
- result.append(confounders)
- return result
-
- def _coerce_causal_model(self, causal_model_cls: Any, data: pd.DataFrame, treatment: str, outcome: str, graph: str):
- try:
- return causal_model_cls(df=data, treatment=treatment, outcome=outcome, graph=graph)
- except TypeError:
- return causal_model_cls(data=data, treatment=treatment, outcome=outcome, graph=graph)
-
- def _parse_p_value(self, value: Any) -> Optional[float]:
- if value is None:
- return None
- if isinstance(value, (int, float, np.floating)):
- return float(value)
- if isinstance(value, (list, tuple)) and value:
- return self._parse_p_value(value[0])
- if isinstance(value, dict):
- for key in ("p_value", "p-value", "p value"):
- if key in value:
- return self._parse_p_value(value[key])
- return None
- match = re.search(r"p[-_ ]?value[^0-9<>=-]*[<>=: ]+\s*([0-9]*\.?[0-9]+)", str(value), re.I)
- if match:
- return float(match.group(1))
- return None
-
- def _as_optional_float(self, value: Any) -> Optional[float]:
- if value is None:
- return None
- try:
- arr = np.asarray(value, dtype=float)
- if arr.size == 1:
- return float(arr.reshape(-1)[0])
- except Exception:
- pass
- try:
- return float(value)
- except Exception:
- return None
-
- def _parse_refuter_result(self, method: str, refute: Any, alpha: float = 0.05) -> Dict[str, object]:
- text = str(refute)
- result_attr = getattr(refute, "refutation_result", None)
- estimated_effect = getattr(refute, "estimated_effect", None)
- new_effect = getattr(refute, "new_effect", None)
- p_value = self._parse_p_value(result_attr)
- if p_value is None:
- p_value = self._parse_p_value(text)
- lower_text = text.lower()
-
- if method == "placebo_treatment":
- if p_value is not None:
- falsified = p_value < alpha
- elif "not statistically significant" in lower_text:
- falsified = False
- elif "statistically significant" in lower_text:
- falsified = True
- else:
- falsified = False
- else:
- if p_value is not None:
- falsified = p_value < alpha
- elif "not statistically significant" in lower_text:
- falsified = False
- elif "statistically significant" in lower_text:
- falsified = True
- else:
- falsified = False
-
- return {
- "method": method,
- "refute": text,
- "estimated_effect": self._as_optional_float(estimated_effect),
- "new_effect": self._as_optional_float(new_effect),
- "p_value": p_value,
- "passed": not falsified,
- "falsified": falsified,
- }
-
- def _run_dowhy_validation(
- self,
- causal_model_cls: Any,
- treatment: str,
- outcome: str,
- treatment_value: float = 1.0,
- adjustment_candidates: Optional[List[List[str]]] = None,
- negative_controls: Optional[List[str]] = None,
- ) -> Dict[str, object]:
- data = pd.DataFrame(self.scm.data_level, columns=self.scm.nodes)
- graph = self._build_dowhy_graph()
- model = self._coerce_causal_model(causal_model_cls, data, treatment, outcome, graph)
- identified_estimand = model.identify_effect()
- estimate = model.estimate_effect(identified_estimand, method_name="backdoor.linear_regression")
-
- refutations = []
- for method in ["placebo_treatment", "random_common_cause"]:
- try:
- kwargs = {"method_name": method}
- if method == "placebo_treatment":
- kwargs["placebo_type"] = "permute"
- refute = model.refute_estimate(identified_estimand, estimate, **kwargs)
- refutations.append(self._parse_refuter_result(method, refute))
- except TypeError:
- try:
- kwargs = {"method_name": method}
- if method == "placebo_treatment":
- kwargs["placebo_type"] = "permute"
- refute = model.refute_estimate(estimate, **kwargs)
- refutations.append(self._parse_refuter_result(method, refute))
- except Exception as exc:
- refutations.append({"method": method, "error": str(exc), "passed": False, "falsified": False})
- except Exception as exc:
- refutations.append({"method": method, "error": str(exc), "passed": False, "falsified": False})
-
- negative_control_checks = []
- for control in negative_controls or []:
- if control in {treatment, outcome}:
- continue
- try:
- nc_model = self._coerce_causal_model(causal_model_cls, data, treatment, control, graph)
- nc_identified = nc_model.identify_effect()
- nc_estimate = nc_model.estimate_effect(nc_identified, method_name="backdoor.linear_regression")
- negative_control_checks.append(
- {
- "control": control,
- "identified_estimand": str(nc_identified),
- "estimate": str(nc_estimate),
- }
- )
- except Exception as exc:
- negative_control_checks.append({"control": control, "error": str(exc)})
-
- falsified = any(bool(item.get("falsified")) for item in refutations)
- return {
- "available": True,
- "falsified": falsified,
- "identified_estimand": str(identified_estimand),
- "estimate": str(estimate),
- "refutations": refutations,
- "adjustment_candidates": adjustment_candidates or [],
- "negative_control_checks": negative_control_checks,
- }
-
- def validate_with_dowhy(
- self,
- treatment: str,
- outcome: str,
- treatment_value: float = 1.0,
- num_placebo: int = 5,
- ) -> Dict[str, object]:
- try:
- from dowhy import CausalModel
- except Exception as exc:
- return {
- "available": False,
- "reason": str(exc),
- "falsified": False,
- "summary": "DoWhy is not installed or failed to import.",
- }
-
- try:
- return self._run_dowhy_validation(
- CausalModel,
- treatment=treatment,
- outcome=outcome,
- treatment_value=treatment_value,
- )
- except Exception as exc:
- return {
- "available": False,
- "reason": str(exc),
- "falsified": False,
- "summary": "DoWhy refutation failed.",
- }
-
- def validate_with_pywhyllm_and_dowhy(
- self,
- treatment: str,
- outcome: str,
- treatment_value: float = 1.0,
- max_edges: Optional[int] = None,
- ) -> Dict[str, object]:
- warnings: List[str] = []
- pywhyllm_report = self.analyze_assumptions_with_pywhyllm(treatment, outcome, max_edges=max_edges)
- adjustment_candidates = self._valid_backdoor_sets(pywhyllm_report, treatment, outcome)
- negative_controls = [
- node
- for node in self._valid_nodes(pywhyllm_report.get("negative_controls") or [])
- if node not in {treatment, outcome}
- ]
-
- try:
- from dowhy import CausalModel
- except Exception as exc:
- return {
- "pywhyllm": pywhyllm_report,
- "dowhy": {
- "available": False,
- "reason": str(exc),
- "falsified": False,
- "summary": "DoWhy is not installed or failed to import.",
- "adjustment_candidates": adjustment_candidates,
- "negative_controls": negative_controls,
- },
- "falsified": False,
- "warnings": warnings,
- }
-
- try:
- dowhy_report = self._run_dowhy_validation(
- CausalModel,
- treatment=treatment,
- outcome=outcome,
- treatment_value=treatment_value,
- adjustment_candidates=adjustment_candidates,
- negative_controls=negative_controls,
- )
- except Exception as exc:
- dowhy_report = {
- "available": False,
- "reason": str(exc),
- "falsified": False,
- "summary": "DoWhy refutation failed.",
- "adjustment_candidates": adjustment_candidates,
- "negative_controls": negative_controls,
- }
-
- warnings.extend(pywhyllm_report.get("warnings") or [])
- return {
- "pywhyllm": pywhyllm_report,
- "dowhy": dowhy_report,
- "falsified": bool(dowhy_report.get("falsified")),
- "warnings": warnings,
- }
-
- def counterfactual(
- self,
- observed_t: int,
- treatment: Optional[str] = None,
- cf_value: Optional[float] = None,
- target: str = "",
- treatments: Optional[Dict[str, float]] = None,
- ) -> Dict[str, object]:
- t = observed_t if observed_t >= 0 else (self.scm.t_steps + observed_t)
- if t < 0 or t >= self.scm.t_steps:
- raise ValueError(f"observed_t {observed_t} resolves out of bounds for T={self.scm.t_steps}")
-
- if treatments is None:
- if treatment is None or cf_value is None:
- raise ValueError("Either treatment/cf_value or treatments must be provided.")
- treatments = {treatment: cf_value}
- elif treatment is not None or cf_value is not None:
- raise ValueError("Provide either treatment/cf_value or treatments, not both.")
-
- observed = {node: float(self.scm.data_level[t, i]) for i, node in enumerate(self.scm.nodes)}
- result = self._counterfactual_outcome(observed, treatments, target)
- result["explanation"] = (
- f"Counterfactual computed at t={t}: set {treatments} "
- f"and propagated structural equations with abducted residuals."
- )
- if len(treatments) > 1:
- result["shapley_contributions"] = self._shapley_contributions(observed, treatments, target)
- else:
- result["shapley_contributions"] = {next(iter(treatments)): result["ite"]}
- result["ite_total"] = result["ite"]
- return result
diff --git a/singular_ticker_causal/causal_inference/tests/test_causal_queries.py b/singular_ticker_causal/causal_inference/tests/test_causal_queries.py
deleted file mode 100644
index 27bc9258f66d5ccc3dc602a344672aac327c9e33..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/causal_inference/tests/test_causal_queries.py
+++ /dev/null
@@ -1,102 +0,0 @@
-import importlib.util
-import numpy as np
-
-from causal_inference.causal_model import StructuralCausalModel
-from causal_inference.query_engine import CausalQueryEngine
-
-
-def _synthetic_inputs():
- nodes = ["A", "B", "C"]
- T = 40
- data = np.zeros((T, 3, 1), dtype=float)
- rng = np.random.default_rng(42)
-
- a = rng.normal(0, 1, size=T)
- b = 0.6 * np.roll(a, 1) + rng.normal(0, 0.1, size=T)
- c = 0.7 * np.roll(b, 1) + rng.normal(0, 0.1, size=T)
- b[0] = rng.normal()
- c[0] = rng.normal()
- data[:, 0, 0] = a
- data[:, 1, 0] = b
- data[:, 2, 0] = c
-
- # A->B, B->C and a weak cycle C->A to exercise pruning.
- adj = np.array(
- [
- [0.0, 0.9, 0.0],
- [0.0, 0.0, 0.8],
- [0.1, 0.0, 0.0],
- ],
- dtype=float,
- )
- prior = np.zeros((3, 3), dtype=float)
- mask = np.ones((T, 3), dtype=float)
- return nodes, adj, prior, data, mask
-
-
-def test_scm_fit_and_cycle_prune():
- nodes, adj, prior, data, mask = _synthetic_inputs()
- scm = StructuralCausalModel(nodes=nodes, adj=adj, adjacency_mask=prior, prohibition_mask=prior, data_tech=data, mask_tech=mask).fit()
- assert scm.dag_adj.shape == (3, 3)
- assert len(scm.topological_indices) == 3
- assert np.sum(scm.dag_adj) <= 2
-
-
-def test_intervention_and_counterfactual_shapes():
- nodes, adj, prior, data, mask = _synthetic_inputs()
- scm = StructuralCausalModel(nodes=nodes, adj=adj, adjacency_mask=prior, prohibition_mask=prior, data_tech=data, mask_tech=mask).fit()
- engine = CausalQueryEngine(scm)
-
- inter = engine.intervene("A", value=0.5, targets=["B", "C"], horizon=3)
- assert set(inter["predicted_values"].keys()) == {"B", "C"}
- assert len(inter["predicted_values"]["B"]) == 3
-
- cf = engine.counterfactual(observed_t=-1, treatment="A", cf_value=0.8, target="C")
- assert "factual_outcome" in cf
- assert "counterfactual_outcome" in cf
- assert "ite" in cf
-
-
-def test_counterfactual_shapley_contributions():
- nodes, adj, prior, data, mask = _synthetic_inputs()
- scm = StructuralCausalModel(nodes=nodes, adj=adj, adjacency_mask=prior, prohibition_mask=prior, data_tech=data, mask_tech=mask).fit()
- engine = CausalQueryEngine(scm)
-
- cf = engine.counterfactual(
- observed_t=-1,
- treatments={"A": 0.8, "B": -0.3},
- target="C",
- )
-
- assert "shapley_contributions" in cf
- assert set(cf["shapley_contributions"].keys()) == {"A", "B"}
- assert abs(sum(cf["shapley_contributions"].values()) - cf["ite"]) < 1e-6
-
-
-def test_dowhy_validation_is_optional():
- nodes, adj, prior, data, mask = _synthetic_inputs()
- scm = StructuralCausalModel(nodes=nodes, adj=adj, adjacency_mask=prior, prohibition_mask=prior, data_tech=data, mask_tech=mask).fit()
- engine = CausalQueryEngine(scm)
-
- result = engine.validate_with_dowhy(treatment="A", outcome="C")
- assert "available" in result
- if importlib.util.find_spec("dowhy") is None:
- assert result["available"] is False
- else:
- assert "falsified" in result
-
-
-def test_exogenous_fallback_prevention():
- nodes, adj, prior, data, mask = _synthetic_inputs()
- # Force fallback to exogenous by setting min_obs to a very large number (e.g. 100)
- scm = StructuralCausalModel(nodes=nodes, adj=adj, adjacency_mask=prior, prohibition_mask=prior, data_tech=data, mask_tech=mask, min_obs=100).fit()
- engine = CausalQueryEngine(scm)
-
- inter = engine.intervene("A", value=0.5, targets=["B", "C"], horizon=3)
- assert set(inter["predicted_values"].keys()) == {"B", "C"}
- assert len(inter["predicted_values"]["B"]) == 3
-
- cf = engine.counterfactual(observed_t=-1, treatment="A", cf_value=0.8, target="C")
- assert "factual_outcome" in cf
- assert "counterfactual_outcome" in cf
- assert "ite" in cf
diff --git a/singular_ticker_causal/causal_inference/tests/test_pywhyllm_assumptions.py b/singular_ticker_causal/causal_inference/tests/test_pywhyllm_assumptions.py
deleted file mode 100644
index 10bf3c056bfc1cf01592c29f463ea8bc30389a82..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/causal_inference/tests/test_pywhyllm_assumptions.py
+++ /dev/null
@@ -1,193 +0,0 @@
-import builtins
-import sys
-import types
-
-import numpy as np
-
-from singular_ticker_causal.causal_inference.causal_model import StructuralCausalModel
-from singular_ticker_causal.causal_inference.pywhyllm_assumptions import (
- PyWhyLLMConfig,
- PyWhyLLMAssumptionService,
-)
-from singular_ticker_causal.causal_inference.query_engine import CausalQueryEngine
-
-
-def _synthetic_scm():
- nodes = ["A", "B", "C"]
- data = np.zeros((24, 3, 1), dtype=float)
- rng = np.random.default_rng(11)
- data[:, 0, 0] = rng.normal(size=24)
- data[:, 1, 0] = rng.normal(size=24)
- data[:, 2, 0] = rng.normal(size=24)
- adj = np.array(
- [
- [0.0, 0.8, 0.7],
- [0.0, 0.0, 0.9],
- [0.0, 0.0, 0.0],
- ],
- dtype=float,
- )
- prior = np.ones((3, 3), dtype=float)
- np.fill_diagonal(prior, 0.0)
- return StructuralCausalModel(
- nodes=nodes,
- adj=adj,
- adjacency_mask=np.zeros((3, 3), dtype=float),
- prohibition_mask=prior,
- data_tech=data,
- ).fit()
-
-
-class FakeModelSuggester:
- def suggest_domain_expertises(self, all_factors):
- return ["financial accounting"]
-
- def suggest_confounders(self, treatment, outcome, all_factors, domain_expertises):
- return ({("B", treatment): 1, ("B", outcome): 1}, ["B", "Missing_Node"])
-
- def suggest_relationships(self, treatment, outcome, all_factors, domain_expertises, strategy):
- return [(treatment, outcome), ("B", outcome)]
-
-
-class FakeIdentificationSuggester:
- def suggest_backdoor(self, treatment, outcome, all_factors, domain_expertises):
- return ["B", "Missing_Node"]
-
- def suggest_mediators(self, treatment, outcome, all_factors, domain_expertises):
- return ["B"]
-
- def suggest_ivs(self, treatment, outcome, all_factors, domain_expertises):
- return ["A"]
-
-
-class FakeValidationSuggester:
- def critique_graph(self, all_factors, suggested_dag, domain_expertises, strategy):
- return "A -> C accepted; B -> C implausible, reject"
-
- def suggest_latent_confounders(self, treatment, outcome, all_factors, domain_expertises):
- return ["market regime"]
-
- def suggest_negative_controls(self, treatment, outcome, all_factors, domain_expertises):
- return ["B", "Missing_Node"]
-
-
-def _fake_service(tmp_path):
- return PyWhyLLMAssumptionService(
- PyWhyLLMConfig(enabled=True, cache_dir=str(tmp_path)),
- model_suggester=FakeModelSuggester(),
- identification_suggester=FakeIdentificationSuggester(),
- validation_suggester=FakeValidationSuggester(),
- relationship_strategy="pairwise",
- )
-
-
-def test_pywhyllm_assumption_service_uses_fakes_and_cache(tmp_path):
- scm = _synthetic_scm()
- service = _fake_service(tmp_path)
-
- report = service.analyze(nodes=scm.nodes, dag_adj=scm.dag_adj, treatment="A", outcome="C")
-
- assert report.available is True
- assert report.domain_expertises == ["financial accounting"]
- assert report.suggested_confounders == ["B", "Missing_Node"]
- assert report.suggested_backdoor_sets == [["B", "Missing_Node"]]
- assert report.negative_controls == ["B", "Missing_Node"]
- assert ("B", "C") in report.rejected_edges
-
- cached = service.analyze(nodes=scm.nodes, dag_adj=scm.dag_adj, treatment="A", outcome="C")
- assert cached.to_dict() == report.to_dict()
-
-
-def test_missing_pywhyllm_returns_unavailable(monkeypatch, tmp_path):
- original_import = builtins.__import__
-
- def fake_import(name, *args, **kwargs):
- if name.startswith("pywhyllm"):
- raise ImportError("blocked pywhyllm")
- return original_import(name, *args, **kwargs)
-
- monkeypatch.setattr(builtins, "__import__", fake_import)
- scm = _synthetic_scm()
- service = PyWhyLLMAssumptionService(PyWhyLLMConfig(enabled=True, cache_dir=str(tmp_path)))
-
- report = service.analyze(nodes=scm.nodes, dag_adj=scm.dag_adj, treatment="A", outcome="C")
-
- assert report.available is False
- assert "blocked pywhyllm" in report.reason
-
-
-def test_engine_filters_pywhyllm_adjustments_and_negative_controls(tmp_path):
- scm = _synthetic_scm()
- engine = CausalQueryEngine(scm, pywhyllm_service=_fake_service(tmp_path), pywhyllm_enabled=True)
-
- report = engine.analyze_assumptions_with_pywhyllm("A", "C")
- adjustment_sets = engine._valid_backdoor_sets(report, "A", "C")
- negative_controls = engine._valid_nodes(report["negative_controls"])
-
- assert adjustment_sets == [["B"]]
- assert negative_controls == ["B"]
-
-
-def test_placebo_not_statistically_significant_passes():
- scm = _synthetic_scm()
- engine = CausalQueryEngine(scm)
-
- parsed = engine._parse_refuter_result(
- "placebo_treatment",
- "Refute: Use a Placebo Treatment. The result is not statistically significant.",
- )
-
- assert parsed["passed"] is True
- assert parsed["falsified"] is False
-
-
-def test_combined_validation_reports_unavailable_dowhy(monkeypatch, tmp_path):
- monkeypatch.setitem(sys.modules, "dowhy", None)
- scm = _synthetic_scm()
- engine = CausalQueryEngine(scm, pywhyllm_service=_fake_service(tmp_path), pywhyllm_enabled=True)
-
- result = engine.validate_with_pywhyllm_and_dowhy("A", "C")
-
- assert result["pywhyllm"]["available"] is True
- assert result["dowhy"]["available"] is False
- assert result["dowhy"]["adjustment_candidates"] == [["B"]]
-
-
-def test_combined_validation_uses_fake_dowhy(monkeypatch, tmp_path):
- class FakeRefute:
- estimated_effect = 1.0
- new_effect = 0.0
- refutation_result = {"p_value": 0.8}
-
- def __str__(self):
- return "not statistically significant"
-
- class FakeEstimate:
- def __str__(self):
- return "estimate"
-
- class FakeCausalModel:
- def __init__(self, **kwargs):
- self.kwargs = kwargs
-
- def identify_effect(self):
- return "estimand"
-
- def estimate_effect(self, identified_estimand, method_name):
- return FakeEstimate()
-
- def refute_estimate(self, identified_estimand, estimate, **kwargs):
- return FakeRefute()
-
- fake_dowhy = types.ModuleType("dowhy")
- fake_dowhy.CausalModel = FakeCausalModel
- monkeypatch.setitem(sys.modules, "dowhy", fake_dowhy)
-
- scm = _synthetic_scm()
- engine = CausalQueryEngine(scm, pywhyllm_service=_fake_service(tmp_path), pywhyllm_enabled=True)
-
- result = engine.validate_with_pywhyllm_and_dowhy("A", "C")
-
- assert result["dowhy"]["available"] is True
- assert result["dowhy"]["falsified"] is False
- assert result["dowhy"]["negative_control_checks"][0]["control"] == "B"
diff --git a/singular_ticker_causal/data_sources/__init__.py b/singular_ticker_causal/data_sources/__init__.py
deleted file mode 100644
index 6fd7c26e5c9519a9dfbae30e339d331bd62386ab..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/data_sources/__init__.py
+++ /dev/null
@@ -1,4 +0,0 @@
-from .fetcher import Fetcher
-from .gdelt_client import GDELTClient
-from .news_client import NewsClient
-from .sebi_reg30_client import SEBIREG30Client
\ No newline at end of file
diff --git a/singular_ticker_causal/data_sources/bsedata/__init__.py b/singular_ticker_causal/data_sources/bsedata/__init__.py
deleted file mode 100644
index d1fd6cb802d787897a4d2a8338e52b089c81c59f..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/data_sources/bsedata/__init__.py
+++ /dev/null
@@ -1,27 +0,0 @@
-"""
-
- MIT License
-
- Copyright (c) 2018 - 2024 Shrey Dabhi
-
- Permission is hereby granted, free of charge, to any person obtaining a copy
- of this software and associated documentation files (the "Software"), to deal
- in the Software without restriction, including without limitation the rights
- to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
- copies of the Software, and to permit persons to whom the Software is
- furnished to do so, subject to the following conditions:
-
- The above copyright notice and this permission notice shall be included in all
- copies or substantial portions of the Software.
-
- THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
- IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
- FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
- AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
- LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
- OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
- SOFTWARE.
-
-"""
-
-name = "bsedata"
\ No newline at end of file
diff --git a/singular_ticker_causal/data_sources/bsedata/bhavcopy.py b/singular_ticker_causal/data_sources/bsedata/bhavcopy.py
deleted file mode 100644
index 97bdb1d3e25f11d238a2b84171a66ee09012c08e..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/data_sources/bsedata/bhavcopy.py
+++ /dev/null
@@ -1,58 +0,0 @@
-import os
-import io
-import csv
-import requests
-import tempfile
-import datetime
-from zipfile import ZipFile
-from .exceptions import BhavCopyNotFound
-from .helpers import COMMON_REQUEST_HEADERS
-
-
-def loadBhavCopyData(statsDate: datetime.date) -> list:
- tempDir = os.path.join(tempfile.gettempdir(), "bsedata")
- zipfileName = f"EQ{statsDate.strftime('%d%m%y')}_CSV.ZIP"
- r = requests.get(
- f"https://www.bseindia.com/download/BhavCopy/Equity/{zipfileName}",
- headers=COMMON_REQUEST_HEADERS,
- )
-
- if r.status_code != 200:
- raise BhavCopyNotFound()
-
- try:
- os.makedirs(tempDir)
- except FileExistsError:
- pass
-
- f_zip = open(os.path.join(tempDir, zipfileName), "wb+")
- f_zip.write(r.content)
- f_zip.close()
-
- output = []
-
- with ZipFile(os.path.join(tempDir, zipfileName)) as bhavCopyZip:
- with bhavCopyZip.open(f"EQ{statsDate.strftime('%d%m%y')}.CSV") as bhavCopyFile:
- reader = csv.DictReader(io.TextIOWrapper(bhavCopyFile))
- for row in reader:
- output.append(mapBhavCopyRowToDict(row))
-
- return output
-
-
-def mapBhavCopyRowToDict(row: dict) -> dict:
- SC_TYPE_MAP = {"B": "bond", "Q": "equity", "D": "debenture", "P": "preference"}
- return {
- "scripCode": row["SC_CODE"],
- "open": row["OPEN"],
- "high": row["HIGH"],
- "low": row["LOW"],
- "close": row["CLOSE"],
- "last": row["LAST"],
- "prevClose": row["PREVCLOSE"],
- "totalTrades": row["NO_TRADES"],
- "totalSharesTraded": row["NO_OF_SHRS"],
- "netTurnover": row["NET_TURNOV"],
- "scripType": SC_TYPE_MAP[row["SC_TYPE"]],
- "securityID": row["SC_NAME"].strip(),
- }
diff --git a/singular_ticker_causal/data_sources/bsedata/bse.py b/singular_ticker_causal/data_sources/bsedata/bse.py
deleted file mode 100644
index 2d3a6d473f167ffd635f59f6db1b268017e4fd32..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/data_sources/bsedata/bse.py
+++ /dev/null
@@ -1,150 +0,0 @@
-"""
-
- MIT License
-
- Copyright (c) 2018 - 2024 Shrey Dabhi
-
- Permission is hereby granted, free of charge, to any person obtaining a copy
- of this software and associated documentation files (the "Software"), to deal
- in the Software without restriction, including without limitation the rights
- to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
- copies of the Software, and to permit persons to whom the Software is
- furnished to do so, subject to the following conditions:
-
- The above copyright notice and this permission notice shall be included in all
- copies or substantial portions of the Software.
-
- THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
- IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
- FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
- AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
- LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
- OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
- SOFTWARE.
-
-"""
-
-from .gainers import getGainers
-from .losers import getLosers
-from .bhavcopy import loadBhavCopyData
-from .quote import quote
-from .indices import indices
-import datetime
-import requests
-import json
-
-
-class BSE(object):
- """
- Class which implements the functionality for
- Bombay Stock Exchange (BSE)
- """
-
- def __init__(self, update_codes=False):
- self.__update_codes = update_codes
- if update_codes:
- self.updateScripCodes()
-
- def topGainers(self):
- """
- :returns: A sorted list of codes of top gainers
- """
- return getGainers()
-
- def topLosers(self):
- """
- :returns: A sorted list of codes of top losers
- """
- return getLosers()
-
- def getQuote(self, scripCode):
- """
- :param scripCode: A stock code
- :returns: A dictionary which contain details about the stock
- :raises InvalidStockException: Raised for stocks which have been suspended or no longer trading on BSE
- """
- return quote(scripCode)
-
- def getIndices(self, category):
- """
- :param category: A category of indices
- :returns: A dictionary with details about the indices belonging to the given category
- """
- return indices(category)
-
- def updateScripCodes(self):
- """
- Download a fresh copy of the scrip code listing
-
- :returns: None
- """
- r = requests.get("https://pub-87b187a07d9c42109c9e6999439a583f.r2.dev/stk.json")
- f_stk = open("stk.json", "w+")
- f_stk.write(json.dumps(r.json()))
- f_stk.close()
- return
-
- def getBhavCopyData(self, statsDate: datetime.date):
- """
- Get historical OHLCV data from Bhav Copy released by BSE everyday after market closing.
- The columns available in the data and their description is as given below.
-
- .. list-table::
- :widths: 25 75
- :header-rows: 1
-
- * - Dictionary Field
- - Description
- * - scripCode
- - Unique code assigned to a scrip of a company by BSE
- * - open
- - The price at which the security first trades on a given trading day
- * - high
- - The highest intra-day price of a stock
- * - low
- - The lowest intra-day price of a stock
- * - close
- - The final price at which a security is traded on a given trading day
- * - last
- - The last trade price of the stock
- * - prevClose
- - The closing price of the stock for the previous trading day
- * - totalTrades
- - The total number of trades of a scrip
- * - totalSharesTraded
- - The total number of shares transacted of a scrip
- * - netTurnover
- - Total turnover of a scrip
- * - scripType
- - Scrip category: Equity, Preference, Debenture or Bond
- * - securityID
- - Name of the company
-
- The Bhav Copy files have been mapped to the above mentioned custom fields. The complete documentation for Bhav Copy can be found here: https://www.bseindia.com/markets/MarketInfo/BhavCopy.aspx.
-
-
- :param statsDate: A `datetime.date` object for the for which you want to fetch the data
- :returns: A list of dictionaries which contains OHLCV data for that day for all scrip codes active on that day
- :raises BhavCopyNotFound: Raised when Bhav Copy file is not found on BSE
- """
- return loadBhavCopyData(statsDate)
-
- def getScripCodes(self):
- """
- :returns: A dictionary with scrip codes as keys and company names as values
- """
- f = open("stk.json", "r")
- return json.loads(f.read())
-
- def verifyScripCode(self, code):
- """
- :returns: Company name if it is a valid stock code, else None
- """
- data = self.getScripCodes()
- return data.get(code)
-
- def __str__(self):
- return "Driver Class for Bombay Stock Exchange (BSE)"
-
- def __repr__(self):
- return f"<{self.__class__.__name__}: update_codes={self.__update_codes}> Driver Class for Bombay Stock Exchange (BSE)"
diff --git a/singular_ticker_causal/data_sources/bsedata/exceptions.py b/singular_ticker_causal/data_sources/bsedata/exceptions.py
deleted file mode 100644
index 6fa37865467c29247050bbb823058c4bbb3e731c..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/data_sources/bsedata/exceptions.py
+++ /dev/null
@@ -1,51 +0,0 @@
-"""
-
- MIT License
-
- Copyright (c) 2018 - 2024 Shrey Dabhi
-
- Permission is hereby granted, free of charge, to any person obtaining a copy
- of this software and associated documentation files (the "Software"), to deal
- in the Software without restriction, including without limitation the rights
- to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
- copies of the Software, and to permit persons to whom the Software is
- furnished to do so, subject to the following conditions:
-
- The above copyright notice and this permission notice shall be included in all
- copies or substantial portions of the Software.
-
- THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
- IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
- FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
- AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
- LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
- OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
- SOFTWARE.
-
-"""
-
-
-class InvalidStockException(Exception):
- """
- Exception raised for stocks which have been suspended or no longer trading on BSE.
-
- :param status: the status of the stock as mentioned on BSE website
- """
-
- def __init__(self, status: str = "Inactive stock"):
- if status == "":
- self.status = "Inactive stock"
- else:
- self.status = status
- super().__init__(self.status)
-
-
-class BhavCopyNotFound(Exception):
- """
- Exception raised when the BhavCopy file is not found on BSE website.
- """
-
- def __init__(self):
- super().__init__(
- """The BhavCopy file was not found on the BSE website. You are probably trying to get data for a trading holiday."""
- )
diff --git a/singular_ticker_causal/data_sources/bsedata/gainers.py b/singular_ticker_causal/data_sources/bsedata/gainers.py
deleted file mode 100644
index 2efe89797c6089e975bc5e1a11ad72c569d598b7..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/data_sources/bsedata/gainers.py
+++ /dev/null
@@ -1,57 +0,0 @@
-"""
-
- MIT License
-
- Copyright (c) 2018 - 2024 Shrey Dabhi
-
- Permission is hereby granted, free of charge, to any person obtaining a copy
- of this software and associated documentation files (the "Software"), to deal
- in the Software without restriction, including without limitation the rights
- to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
- copies of the Software, and to permit persons to whom the Software is
- furnished to do so, subject to the following conditions:
-
- The above copyright notice and this permission notice shall be included in all
- copies or substantial portions of the Software.
-
- THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
- IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
- FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
- AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
- LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
- OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
- SOFTWARE.
-
-"""
-
-from .helpers import COMMON_REQUEST_HEADERS
-from bs4 import BeautifulSoup as bs
-import requests
-
-
-def getGainers() -> dict:
- baseurl = """https://m.bseindia.com"""
- res = requests.get(baseurl, headers=COMMON_REQUEST_HEADERS)
- c = res.content
- soup = bs(c, "lxml")
- for tag in soup("div"):
- try:
- if tag["id"] == "divGainers":
- resSoup = tag
- break
- except KeyError:
- continue
- children = list(resSoup.table.contents)
- children = children[1:]
- gainers = []
- for tr in children:
- td = tr.contents
- gainer = {
- "securityID": str(td[0].a.string),
- "scripCode": str(tr.td.a["href"].split("=")[1]),
- "LTP": str(td[1].string),
- "change": str(td[2].string),
- "pChange": str(td[3].string),
- }
- gainers.append(gainer)
- return gainers
diff --git a/singular_ticker_causal/data_sources/bsedata/helpers.py b/singular_ticker_causal/data_sources/bsedata/helpers.py
deleted file mode 100644
index 1e679690cfad3d41f210d961f7ff27e8769572bb..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/data_sources/bsedata/helpers.py
+++ /dev/null
@@ -1,3 +0,0 @@
-COMMON_REQUEST_HEADERS = {
- "User-Agent": "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/83.0.4103.97 Safari/537.36 Edg/83.0.478.45"
-}
diff --git a/singular_ticker_causal/data_sources/bsedata/indices.py b/singular_ticker_causal/data_sources/bsedata/indices.py
deleted file mode 100644
index f3f9982a3ea554bd4e7a6c369c13747f1a881036..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/data_sources/bsedata/indices.py
+++ /dev/null
@@ -1,112 +0,0 @@
-"""
-
- MIT License
-
- Copyright (c) 2018 - 2024 Shrey Dabhi
-
- Permission is hereby granted, free of charge, to any person obtaining a copy
- of this software and associated documentation files (the "Software"), to deal
- in the Software without restriction, including without limitation the rights
- to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
- copies of the Software, and to permit persons to whom the Software is
- furnished to do so, subject to the following conditions:
-
- The above copyright notice and this permission notice shall be included in all
- copies or substantial portions of the Software.
-
- THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
- IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
- FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
- AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
- LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
- OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
- SOFTWARE.
-
-"""
-
-from .helpers import COMMON_REQUEST_HEADERS
-from bs4 import BeautifulSoup as bs
-import requests
-
-
-def indices(category: str) -> dict:
- cat = {
- "market_cap/broad": "1,2",
- "sector_and_industry": "2,2",
- "thematics": "3,2",
- "strategy": "4,2",
- "sustainability": "5,2",
- "volatility": "6,1",
- "composite": "7,1",
- "government": "8,1",
- "corporate": "9,1",
- "money_market": "10,1",
- }
- try:
- ddl_category = cat[category]
- except KeyError:
- print(
- """
-### Invalid category ###
-Use one of the categories mentioned below:
-
-market_cap/broad
-sector_and_industry
-thematics
-strategy
-sustainability
-volatility
-composite
-government
-corporate
-money_market
- """
- )
- return
- baseurl = """https://m.bseindia.com/IndicesView_New.aspx"""
- res = requests.get(baseurl, headers=COMMON_REQUEST_HEADERS)
- c = res.content
- soup = bs(c, "lxml")
- options = {
- "__EVENTTARGET": "ddl_Category",
- "__VIEWSTATEENCRYPTED": "",
- "__EVENTARGUMENT": "",
- "__LASTFOCUS": "",
- "__VIEWSTATEGENERATOR": "162C96CD",
- "UcHeaderMenu1$txtGetQuote": "",
- "__EVENTVALIDATION": "",
- "__VIEWSTATE": "",
- }
- for input in soup("input"):
- try:
- if input["type"] == "hidden":
- if input["id"] == "__VIEWSTATE":
- options["__VIEWSTATE"] = input["value"]
- elif input["id"] == "__EVENTVALIDATION":
- options["__EVENTVALIDATION"] = input["value"]
- except KeyError:
- continue
- options["ddl_Category"] = ddl_category
- res = requests.post(url=baseurl, data=options, headers=COMMON_REQUEST_HEADERS)
- c = res.content
- soup = bs(c, "lxml")
- index_list = []
- for td in soup("td"):
- try:
- if td["class"][0] == "TTRow_left":
- index = {}
- index["currentValue"] = td.next_sibling.string.strip()
- index["change"] = td.next_sibling.next_sibling.string.strip()
- index[
- "pChange"
- ] = td.next_sibling.next_sibling.next_sibling.string.strip()
- index["scripFlag"] = td.a["href"].strip().split("=")[1]
- index["name"] = td.a.string.strip().replace(";", "")
- index_list.append(index)
- except KeyError:
- continue
- results = {}
- for span in soup("span", id="inddate"):
- results["updatedOn"] = span.string[6:].split("|")[0].strip()
- results["indices"] = index_list
- return results
diff --git a/singular_ticker_causal/data_sources/bsedata/losers.py b/singular_ticker_causal/data_sources/bsedata/losers.py
deleted file mode 100644
index c4f895e3f1df1f23127f2ea75825a9e188f01cd8..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/data_sources/bsedata/losers.py
+++ /dev/null
@@ -1,57 +0,0 @@
-"""
-
- MIT License
-
- Copyright (c) 2018 - 2024 Shrey Dabhi
-
- Permission is hereby granted, free of charge, to any person obtaining a copy
- of this software and associated documentation files (the "Software"), to deal
- in the Software without restriction, including without limitation the rights
- to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
- copies of the Software, and to permit persons to whom the Software is
- furnished to do so, subject to the following conditions:
-
- The above copyright notice and this permission notice shall be included in all
- copies or substantial portions of the Software.
-
- THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
- IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
- FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
- AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
- LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
- OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
- SOFTWARE.
-
-"""
-
-from .helpers import COMMON_REQUEST_HEADERS
-from bs4 import BeautifulSoup as bs
-import requests
-
-
-def getLosers() -> dict:
- baseurl = """https://m.bseindia.com"""
- res = requests.get(baseurl, headers=COMMON_REQUEST_HEADERS)
- c = res.content
- soup = bs(c, "lxml")
- for tag in soup("div"):
- try:
- if tag["id"] == "divLosers":
- resSoup = tag
- break
- except KeyError:
- continue
- children = list(resSoup.table.contents)
- children = children[1:]
- losers = []
- for tr in children:
- td = tr.contents
- loser = {
- "securityID": str(td[0].a.string),
- "scripCode": str(tr.td.a["href"].split("=")[1]),
- "LTP": str(td[1].string),
- "change": str(td[2].string),
- "pChange": str(td[3].string),
- }
- losers.append(loser)
- return losers
diff --git a/singular_ticker_causal/data_sources/bsedata/quote.py b/singular_ticker_causal/data_sources/bsedata/quote.py
deleted file mode 100644
index 0e0ceacacc991eb66ca59e0173f8aa8bb30058ec..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/data_sources/bsedata/quote.py
+++ /dev/null
@@ -1,176 +0,0 @@
-"""
-
- MIT License
-
- Copyright (c) 2018 - 2024 Shrey Dabhi
-
- Permission is hereby granted, free of charge, to any person obtaining a copy
- of this software and associated documentation files (the "Software"), to deal
- in the Software without restriction, including without limitation the rights
- to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
- copies of the Software, and to permit persons to whom the Software is
- furnished to do so, subject to the following conditions:
-
- The above copyright notice and this permission notice shall be included in all
- copies or substantial portions of the Software.
-
- THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
- IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
- FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
- AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
- LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
- OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
- SOFTWARE.
-
-"""
-
-from .exceptions import InvalidStockException
-from .helpers import COMMON_REQUEST_HEADERS
-from datetime import datetime as dt
-from bs4 import BeautifulSoup as bs
-import requests
-
-
-def quote(scripCode: str) -> dict:
- baseurl = """https://m.bseindia.com/StockReach.aspx?scripcd="""
- res = requests.get(baseurl + scripCode, headers=COMMON_REQUEST_HEADERS)
- c = res.content
- soup = bs(c, "lxml")
-
- res = {}
-
- for span in soup("span"):
- updt_date = soup.find("span", id="strongDate").text.split("-")[1].strip()
- updt_diff = dt.strptime(updt_date, "%d %b %y | %I:%M %p") - dt.now()
- if updt_diff.days < -7:
- error_text = ""
- error_text_element = soup.find("td", id="tdDispTxt")
- if error_text_element is not None:
- error_text = error_text_element.text
- raise InvalidStockException(status=error_text)
- try:
- if span["class"][0] == "srcovalue":
- try:
- if span["id"] == "spanchangVal":
- res["change"] = span.string.split("(")[0].strip()
- res["pChange"] = span.string.split("(")[1].strip()[:-2]
- except KeyError:
- res["currentValue"] = span.strong.string
- elif span["class"][0] == "companyname":
- res["companyName"] = span.string
- except KeyError:
- try:
- if span["id"] == "lblPBdate":
- try:
- res["priceBand"] = span.string.split(":")[1].strip()
- except AttributeError:
- res["priceBand"] = ""
- elif span["id"] == "strongDate":
- res["updatedOn"] = span.string.split("-")[1].strip()
- except KeyError:
- continue
-
- for td in soup("td"):
- try:
- if td["id"] == "tdCShortName":
- res["securityID"] = td.string.strip()
- elif td["id"] == "tdscripcode":
- res["scripCode"] = td.string.strip()
- elif td["id"] == "tdgroup":
- res["group"] = td.string.strip()
- elif td["id"] == "tdfacevalue":
- res["faceValue"] = td.string.strip()
- elif td["id"] == "tdIndustry":
- res["industry"] = td.string.strip()
- elif td["id"] == "tdpcloseopen":
- res["previousClose"] = td.string.split("/")[0].strip()
- res["previousOpen"] = td.string.split("/")[1].strip()
- elif td["id"] == "tdDHL":
- res["dayHigh"] = td.string.split("/")[0].strip()
- res["dayLow"] = td.string.split("/")[1].strip()
- elif td["id"] == "td52WHL":
- res["52weekHigh"] = td.string.split("/")[0].strip()
- res["52weekLow"] = td.string.split("/")[1].strip()
- elif td["id"] == "tdWAp":
- res["weightedAvgPrice"] = td.string.strip()
- elif td["id"] == "tdTTV":
- res["totalTradedValue"] = td.string.strip() + " Cr."
- elif td["id"] == "tdTTQW":
- res["totalTradedQuantity"] = td.string.split("/")[0].strip() + " Lakh"
- res["2WeekAvgQuantity"] = td.string.split("/")[1].strip() + " Lakh"
- elif td["id"] == "tdMktCapVal":
- res["marketCapFull"] = td.string.split("/")[0].strip() + " Cr."
- res["marketCapFreeFloat"] = td.string.split("/")[1].strip() + " Cr."
- except KeyError:
- continue
-
- if res.get("priceBand", "") != "":
- for tbody in soup("tbody"):
- try:
- if tbody["id"] == "PBtablebody":
- data = tbody.contents[2]
- res["upperPriceBand"] = data.contents[1].string.strip()
- res["lowerPriceBand"] = data.contents[2].string.strip()
- except KeyError:
- continue
-
- buy = {}
- sell = {}
- for td in soup("td"):
- try:
- if td["id"] == "tdBQ1":
- buy["1"] = {
- "quantity": td.string,
- "price": td.next_sibling.next_sibling.string,
- }
- elif td["id"] == "tdBQ2":
- buy["2"] = {
- "quantity": td.string,
- "price": td.next_sibling.next_sibling.string,
- }
- elif td["id"] == "tdBQ3":
- buy["3"] = {
- "quantity": td.string,
- "price": td.next_sibling.next_sibling.string,
- }
- elif td["id"] == "tdBQ4":
- buy["4"] = {
- "quantity": td.string,
- "price": td.next_sibling.next_sibling.string,
- }
- elif td["id"] == "tdBQ5":
- buy["5"] = {
- "quantity": td.string,
- "price": td.next_sibling.next_sibling.string,
- }
- elif td["id"] == "tdSP1":
- sell["1"] = {
- "price": td.string,
- "quantity": td.next_sibling.next_sibling.string,
- }
- elif td["id"] == "tdSP2":
- sell["2"] = {
- "price": td.string,
- "quantity": td.next_sibling.next_sibling.string,
- }
- elif td["id"] == "tdSP3":
- sell["3"] = {
- "price": td.string,
- "quantity": td.next_sibling.next_sibling.string,
- }
- elif td["id"] == "tdSP4":
- sell["4"] = {
- "price": td.string,
- "quantity": td.next_sibling.next_sibling.string,
- }
- elif td["id"] == "tdSP5":
- sell["5"] = {
- "price": td.string,
- "quantity": td.next_sibling.next_sibling.string,
- }
- except KeyError:
- continue
- res["buy"] = buy
- res["sell"] = sell
-
- return res
diff --git a/singular_ticker_causal/data_sources/fetcher.py b/singular_ticker_causal/data_sources/fetcher.py
deleted file mode 100644
index 4c8bdd36126d4f7d3deaeea1e0ba7cb220230e01..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/data_sources/fetcher.py
+++ /dev/null
@@ -1,219 +0,0 @@
-import pandas as pd
-import numpy as np
-import yfinance as yf
-import logging
-from typing import Optional
-from singular_ticker_causal.services.schema import INCOME_STATEMENT_NODES, BALANCE_SHEET_NODES, STRATEGIC_NODES
-IND_AS_NODES = INCOME_STATEMENT_NODES + BALANCE_SHEET_NODES + STRATEGIC_NODES
-from .nseconnect.nse import Nse
-from .bsedata.bse import BSE
-
-
-logger = logging.getLogger(__name__)
-
-
-class Fetcher:
- """
- Tiered fundamental data fetcher.
- Priority: NSE/BSE → XBRL → IndianAPI → yfinance.
- """
- def __init__(self, ticker: str, api_key: Optional[str] = None):
- self.ticker = ticker if ticker.endswith(".NS") or ticker.endswith(".BO") else f"{ticker}.NS"
- self.api_key = api_key
-
- def fetch(self, start: str, end: str) -> pd.DataFrame:
- logger.info(f"Fetching fundamentals for {self.ticker} from {start} to {end}")
-
- df = None
- if self.ticker.endswith(".NS"):
- logger.info("Attempting NSE fetch...")
- df = self._try_nse(start, end)
- elif self.ticker.endswith(".BO"):
- logger.info("Attempting BSE fetch...")
- df = self._try_bse(start, end)
-
- if df is None or df.empty:
- logger.info("Attempting Tier 1: XBRL...")
- df = self._try_xbrl(start, end)
- if df is None or df.empty:
- logger.info("Tier 1 failed or returned empty. Attempting Tier 2: IndianAPI...")
- df = self._try_indianapi(start, end)
- if df is None or df.empty:
- logger.info("Tier 2 failed or returned empty. Attempting Tier 3: yfinance...")
- df = self._try_yfinance(start, end)
-
- if df is not None and not df.empty:
- df = df.sort_index()
- # Forward fill balance sheet items as they are point-in-time and usually stable
- # We do this AFTER mapping in _try_yfinance, so we use mapped names
- bs_cols = ["Total_Assets", "Shareholders_Equity", "Inventory", "Accounts_Payable", "Total_Debt", "PPE", "CWIP", "Intangible_Assets"]
- available_bs_cols = [c for c in bs_cols if c in df.columns]
- if available_bs_cols:
- df[available_bs_cols] = df[available_bs_cols].ffill()
-
- if df is None or df.empty:
- logger.warning(f"No fundamental data found for {self.ticker}")
- return pd.DataFrame(columns=IND_AS_NODES)
-
- logger.info("Deriving strategic nodes...")
- df = self._derive_strategic_nodes(df)
-
- # Ensure all IND_AS_NODES are present
- for node in IND_AS_NODES:
- if node not in df.columns:
- df[node] = np.nan
-
- return df[IND_AS_NODES].sort_index()
-
- def _try_nse(self, start: str, end: str) -> Optional[pd.DataFrame]:
- # Using nseconnect for high-fidelity NSE data
- try:
- # Clean ticker (e.g. RELIANCE.NS -> RELIANCE)
- clean_ticker = self.ticker.split('.')[0]
- nse = Nse()
- logger.info(f"Attempting to fetch NSE data for {clean_ticker}...")
- # Note: nseconnect is primarily for quotes;
- # for full fundamentals we still rely on yfinance or XBRL.
- # We return None here to let it fall back, but the plumbing is now real.
- quote = nse.get_quote(clean_ticker)
- if quote:
- logger.info(f"Successfully connected to NSE for {clean_ticker}")
- return None
- except Exception as e:
- logger.error(f"Error fetching from NSE: {e}")
- return None
-
- def _try_bse(self, start: str, end: str) -> Optional[pd.DataFrame]:
- # Using bsedata for high-fidelity BSE data
- try:
- # TODO: Implement mapping from alphabetic ticker to numeric BSE scrip code
- bse = BSE()
- logger.info(f"Attempting to fetch BSE data for {self.ticker}...")
- # Currently limited to quotes; returning None to fall back to yfinance
- return None
- except Exception as e:
- logger.error(f"Error fetching from BSE: {e}")
- return None
-
- def _try_xbrl(self, start: str, end: str) -> Optional[pd.DataFrame]:
- # Tier 1 extraction via python-xbrl / Arelle is currently in development
- return None
-
- def _try_indianapi(self, start: str, end: str) -> Optional[pd.DataFrame]:
- # Tier 2 integration for IndianAPI.in / FinEdge API is currently in development
- return None
-
- def _try_yfinance(self, start: str, end: str) -> Optional[pd.DataFrame]:
- try:
- t = yf.Ticker(self.ticker)
- q_fin = t.quarterly_financials.T
- q_bs = t.quarterly_balance_sheet.T
- q_cf = t.quarterly_cashflow.T
-
- if q_fin.empty and q_bs.empty and q_cf.empty:
- logger.warning("All yfinance statements (financials, balance_sheet, cashflow) are empty.")
- return None
-
- # Merge all three statements
- logger.info(f"Merging yfinance statements: q_fin={q_fin.shape}, q_bs={q_bs.shape}, q_cf={q_cf.shape}")
- df = pd.concat([q_fin, q_bs, q_cf], axis=1)
-
-
- df = df.loc[:, ~df.columns.duplicated()] # Remove duplicate columns if any
- logger.info(f"Merged shape after removing duplicates: {df.shape}")
-
- # Map yfinance columns to IND_AS_NODES (Simplified mapping for MVP)
- mapping = {
- "Total Revenue": "Revenue",
- "Cost Of Revenue": "COGS",
- "Operating Expense": "Operating_Expenses",
- "Operating Income": "EBIT",
- "EBIT": "EBIT",
- "EBITDA": "EBITDA",
- "Interest Expense": "Interest_Expense",
- "Pretax Income": "EBT",
- "Tax Provision": "Tax_Expense",
- "Net Income": "PAT",
- "Total Assets": "Total_Assets",
- "Stockholders Equity": "Shareholders_Equity",
- "Depreciation And Amortization": "D_A",
- "Inventory": "Inventory",
- "Accounts Payable": "Accounts_Payable",
- "Total Debt": "Total_Debt",
- "Operating Cash Flow": "Operating_Cash_Flow",
- "Capital Expenditure": "Capex",
- "Net PPE": "PPE",
- "Construction In Progress": "CWIP",
- "Goodwill And Other Intangible Assets": "Intangible_Assets",
- }
-
- df = df.rename(columns=mapping)
- df = df.loc[:, ~df.columns.duplicated()] # Remove duplicates after rename
- df.index = pd.to_datetime(df.index)
- return df
- except Exception as e:
- logger.error(f"Error fetching from yfinance: {e}")
- return None
-
-
- def _derive_strategic_nodes(self, df: pd.DataFrame) -> pd.DataFrame:
- """
- Compute Layer 3 nodes and other derived fields.
- """
- df = df.copy()
-
- # Helper to safely get a series from a column that might be a DataFrame
- def get_series(name):
- if name not in df.columns:
- return None
- col = df[name]
- if isinstance(col, pd.DataFrame):
- logger.warning(f"Column '{name}' is a DataFrame with multiple columns: {col.columns.tolist()}. Taking the first.")
- return col.iloc[:, 0]
- return col
-
- # 1. Average Assets & Equity (Rolling 2-period mean)
- assets = get_series("Total_Assets")
- if assets is not None:
- df["Average_Total_Assets"] = assets.rolling(window=2).mean().fillna(assets)
-
- equity = get_series("Shareholders_Equity")
- if equity is not None:
- df["Average_Shareholders_Equity"] = equity.rolling(window=2).mean().fillna(equity)
-
- # 2. Basic derivations
- rev = get_series("Revenue")
- cogs = get_series("COGS")
- if rev is not None and cogs is not None:
- df["Gross_Profit"] = rev - cogs.fillna(0)
-
- ebit = get_series("EBIT")
- da = get_series("D_A")
- if ebit is not None:
- # Proper EBITDA = EBIT + Depreciation & Amortization
- df["EBITDA"] = ebit + da.fillna(0) if da is not None else ebit
-
- pat = get_series("PAT")
- if pat is not None and rev is not None:
- df["Net_Profit_Margin"] = pat / rev.replace(0, np.nan)
-
- avg_assets = get_series("Average_Total_Assets")
- if rev is not None and avg_assets is not None:
- df["Asset_Turnover"] = rev / avg_assets.replace(0, np.nan)
-
- avg_equity = get_series("Average_Shareholders_Equity")
- if avg_assets is not None and avg_equity is not None:
- df["Equity_Multiplier"] = avg_assets / avg_equity.replace(0, np.nan)
-
- npm = get_series("Net_Profit_Margin")
- at = get_series("Asset_Turnover")
- em = get_series("Equity_Multiplier")
- if npm is not None and at is not None and em is not None:
- df["ROE"] = npm * at * em
-
- ocf = get_series("Operating_Cash_Flow")
- capex = get_series("Capex")
- if ocf is not None and capex is not None:
- df["Free_Cash_Flow"] = ocf - capex.abs().fillna(0)
-
- return df
diff --git a/singular_ticker_causal/data_sources/gdelt_client.py b/singular_ticker_causal/data_sources/gdelt_client.py
deleted file mode 100644
index e681ef612cb91738ee1b62aba65b5670bc8bdecb..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/data_sources/gdelt_client.py
+++ /dev/null
@@ -1,400 +0,0 @@
-import time
-import random
-import logging
-import requests
-import yfinance as yf
-from datetime import datetime, timezone
-from typing import List, Dict, Any
-from singular_ticker_causal.utils.llm_client import LLMClient
-
-
-logger = logging.getLogger(__name__)
-
-
-# ==============================================================
-# GDELT DOC 2.0 FETCHER
-# ==============================================================
-
-def _fmt_gdelt_dt(dt: datetime) -> str:
- """
- Convert datetime -> GDELT YYYYMMDDHHMMSS (UTC).
- """
- if dt.tzinfo is not None:
- dt = dt.astimezone(timezone.utc).replace(tzinfo=None)
-
- return dt.strftime("%Y%m%d%H%M%S")
-
-
-def _dedupe_articles(articles: list[dict]) -> list[dict]:
- """
- Deduplicate by URL/title combination.
- """
- seen = set()
- deduped = []
-
- for article in articles:
- key = (
- article.get("link", "").strip(),
- article.get("title", "").strip().lower(),
- )
-
- if key in seen:
- continue
-
- seen.add(key)
- deduped.append(article)
-
- return deduped
-
-
-class GDELTClient:
- """
- Thin wrapper around GDELT Doc API for macro context.
- """
- GDELT_DOC_API: str = "https://api.gdeltproject.org/api/v2/doc/doc"
- GDELT_MAX_RECORDS: int = 250
- GDELT_MAX_WINDOW_DAYS: int = 90
- USER_AGENT: str = (
- "Mozilla/5.0 (Windows NT 10.0; Win64; x64; rv:128.0) "
- "Gecko/20100101 Firefox/128.0"
- )
-
- def __init__(self):
- pass
-
- def fetch(
- self,
- keyword: str,
- from_dt: datetime,
- to_dt: datetime,
- ticker: str | None = None,
- **kwargs,
- ) -> List[Dict[str, Any]]:
- """
- Main entry point for GDELT fetching.
- If 'ticker' is provided, performs a broad generic search using yfinance context.
- """
- if ticker:
- return self.fetch_ticker_context(ticker, from_dt, to_dt, **kwargs)
-
- logger.info(
- "Fetching GDELT news for %s from %s to %s",
- keyword,
- from_dt,
- to_dt,
- )
-
- try:
- articles = self.fetch_gdelt_window(
- keyword=keyword,
- from_dt=from_dt,
- to_dt=to_dt,
- **kwargs,
- )
-
- for article in articles:
- article["credibility_weight"] = 0.5
- article["source_type"] = "gdelt"
-
- return articles
-
- except Exception as e:
- logger.exception("Error fetching from GDELT: %s", e)
- return []
-
- def _extract_search_terms(self, summary: str, company_name: str) -> Dict[str, List[str]]:
- """Extract product, partner, and industry search terms from business summary using LLM."""
- if not summary:
- return {"products": [], "partners": [], "industry": []}
-
- llm = LLMClient()
- prompt = f"""
- Extract major products/services, strategic partners, and industry-related details from the following summary of {company_name}.
-
- Summary: {summary}
-
- Important instructions:
- 1. For partners, include a few words describing the specific connection or nature of the relationship between {company_name} and the partner (e.g., "{company_name} strategic collaboration with [Partner]").
- 2. For industry, extract a few words describing the sector or niche {company_name} is connected to and its specific role/connection (e.g., "{company_name} provides [Services] in the [Industry] sector").
- 3. Keep it to the top 5 most relevant terms for products and partners, and top 3 for industry.
-
- Return a JSON object with:
- "products": ["product1", "product2", ...],
- "partners": ["{company_name} [connection] [partner1]", ...],
- "industry": ["{company_name} [role/connection] [industry]", ...]
-
- Return ONLY valid JSON.
- """
- try:
- result = llm.chat_json([{"role": "user", "content": prompt}])
- return {
- "products": result.get("products", []),
- "partners": result.get("partners", []),
- "industry": result.get("industry", [])
- }
- except Exception as e:
- logger.error(f"[GDELT] LLM term extraction failed: {e}")
- return {"products": [], "partners": [], "industry": []}
-
- def fetch_ticker_context(
- self,
- ticker_symbol: str,
- from_dt: datetime,
- to_dt: datetime,
- **kwargs
- ) -> List[Dict[str, Any]]:
- """
- Generic fetcher that uses yfinance and LLM to build a broad context.
- Sequence: products, partners, company name, company heads.
- """
- logger.info(f"[GDELT] Fetching generic context for ticker: {ticker_symbol}")
-
- try:
- ticker = yf.Ticker(ticker_symbol)
- info = ticker.info
- except Exception as e:
- logger.error(f"[GDELT] yfinance failed for {ticker_symbol}: {e}")
- return []
-
- summary = info.get("longBusinessSummary", "")
- company_name = info.get("longName") or info.get("shortName") or ticker_symbol
- officers = info.get("companyOfficers", [])
-
- # 1. LLM extraction
- terms = self._extract_search_terms(summary, company_name)
- products = terms.get("products", [])
- partners = terms.get("partners", [])
- industry = terms.get("industry", [])
-
- # 2. Officer names
- heads = [o.get("name") for o in officers if o.get("name")]
-
- # 3. Execution sequence
- all_articles = []
- search_plan = [
- ("products", products),
- ("partners", partners),
- ("industry", industry),
- ("company_name", [company_name]),
- ("company_heads", heads)
- ]
-
- for category, keywords in search_plan:
- for kw in keywords:
- if not kw: continue
- logger.info(f"[GDELT] Scraping category '{category}': {kw}")
- results = self.fetch_gdelt_window(
- keyword=kw,
- from_dt=from_dt,
- to_dt=to_dt,
- **kwargs
- )
- all_articles.extend(results)
- # Polite pause to avoid aggressive rate limiting
- time.sleep(random.uniform(2.0, 5.0))
-
- return _dedupe_articles(all_articles)
-
- def fetch_gdelt_window(
- self,
- keyword: str,
- from_dt: datetime,
- to_dt: datetime,
- max_records: int = 250,
- exact_phrase: bool = False,
- source_country: str | None = None,
- source_lang: str | None = None,
- theme: str | None = None,
- domain: str | None = None,
- extra_query: str | None = None,
- ) -> list:
- """
- Fetch historical articles from the GDELT DOC 2.0 API.
- Includes robust retry logic with 5-minute cooldown for timeouts/rate-limits.
- """
- if from_dt >= to_dt:
- raise ValueError("from_dt must be earlier than to_dt")
-
- window_days = (to_dt - from_dt).days
- if window_days > self.GDELT_MAX_WINDOW_DAYS:
- raise ValueError(
- f"GDELT DOC API only supports ~{self.GDELT_MAX_WINDOW_DAYS} days history"
- )
-
- max_records = min(max_records, self.GDELT_MAX_RECORDS)
- query = self._build_gdelt_query(
- keyword=keyword,
- exact_phrase=exact_phrase,
- source_country=source_country,
- source_lang=source_lang,
- theme=theme,
- domain=domain,
- extra_query=extra_query,
- )
-
- params = {
- "query": query,
- "mode": "artlist",
- "format": "json",
- "maxrecords": max_records,
- "sort": "DateDesc",
- "STARTDATETIME": _fmt_gdelt_dt(from_dt),
- "ENDDATETIME": _fmt_gdelt_dt(to_dt),
- }
-
- max_retries = 5
- for attempt in range(max_retries):
- try:
- resp = requests.get(
- self.GDELT_DOC_API,
- params=params,
- timeout=45,
- headers={
- "User-Agent": self.USER_AGENT,
- "Accept": "application/json",
- },
- )
-
- if resp.status_code == 429:
- logger.warning(
- "[GDELT] Rate limited (429). Cooling down for 5 minutes..."
- )
- time.sleep(305)
- continue
-
- resp.raise_for_status()
-
- try:
- data = resp.json()
- except Exception as e:
- logger.error("[GDELT] Invalid JSON response (Status %d): %s", resp.status_code, e)
- # Log the body to see what GDELT is actually returning (likely HTML error)
- body_snippet = resp.text[:500] if resp.text else "[Empty Response]"
- logger.error("[GDELT] Response body snippet: %s", body_snippet)
-
- # If it's an HTML error page, we might be blocked or throttled in a way that doesn't return 429
- if " str:
- """Build a valid GDELT DOC API query string."""
- query_parts = []
- if keyword:
- if exact_phrase and " " in keyword:
- query_parts.append(f'"{keyword}"')
- else:
- query_parts.append(keyword)
-
- if source_country:
- query_parts.append(f"sourcecountry:{source_country.lower()}")
- if source_lang:
- query_parts.append(f"sourcelang:{source_lang.lower()}")
- if theme:
- query_parts.append(f"theme:{theme}")
- if domain:
- query_parts.append(f"domain:{domain}")
- if extra_query:
- query_parts.append(extra_query)
-
- return " ".join(query_parts)
-
-
-if __name__ == "__main__":
- import json
- from datetime import timedelta
-
- # Configure logging
- logging.basicConfig(
- level=logging.INFO,
- format="%(asctime)s - %(name)s - %(levelname)s - %(message)s"
- )
-
- client = GDELTClient()
-
- # Test parameters
- ticker = "RELIANCE.NS"
- end_date = datetime.now(timezone.utc)
- start_date = end_date - timedelta(days=3)
-
- print(f"\n--- Testing GDELT Ticker Context Fetch: {ticker} ---")
- try:
- articles = client.fetch(
- keyword="",
- from_dt=start_date,
- to_dt=end_date,
- ticker=ticker
- )
-
- output_file = "gdelt.json"
- with open(output_file, "w") as f:
- json.dump(articles, f, indent=4)
-
- print(f"Successfully fetched {len(articles)} articles.")
- print(f"Output saved to {output_file}")
-
- except Exception as e:
- print(f"Test failed: {e}")
\ No newline at end of file
diff --git a/singular_ticker_causal/data_sources/news_client.py b/singular_ticker_causal/data_sources/news_client.py
deleted file mode 100644
index 02a0708067d05ea0e1b8edeeaa962a4d03accdf6..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/data_sources/news_client.py
+++ /dev/null
@@ -1,460 +0,0 @@
-"""
-singular_ticker_causal/data_sources/news_client.py
-
-Fetches quarterly RSS news for a single ticker over a multi-year window.
- - Operates on a single ticker (not a universe)
- - Fetches quarterly windows across the full fundamental date range
- - Returns articles structured for DenoisedNewsEncoder consumption
-
-Sources:
- • LiveMint RSS feeds
- • CNBC-TV18 RSS feeds
- • Other RSS feeds (Business Standard, Forbes India, Zee News, Economic Times, etc.)
- • Trading Economics (Selenium — stream, economy, markets, India news)
- • Zerodha Pulse (requests + BeautifulSoup)
-"""
-
-import os
-import json
-import time
-import logging
-import datetime
-import requests
-import feedparser
-import yfinance as yf
-from typing import List
-from bs4 import BeautifulSoup
-from selenium import webdriver
-from selenium.webdriver.firefox.options import Options as FirefoxOptions
-from selenium.webdriver.firefox.service import Service as FirefoxService
-from selenium.webdriver.common.by import By
-from selenium.webdriver.support.ui import WebDriverWait
-from selenium.webdriver.support import expected_conditions as EC
-
-
-logger = logging.getLogger(__name__)
-
-
-# ==============================================================
-# SCRAPER SOURCES
-# ==============================================================
-
-LIVEMINT_FEEDS = [
- "https://www.livemint.com/rss/companies",
- "https://www.livemint.com/rss/opinion",
- "https://www.livemint.com/rss/money",
- "https://www.livemint.com/rss/politics",
- "https://www.livemint.com/rss/science",
- "https://www.livemint.com/rss/industry",
- "https://www.livemint.com/rss/education",
- "https://www.livemint.com/rss/sports",
- "https://www.livemint.com/rss/technology",
- "https://www.livemint.com/rss/news",
- "https://www.livemint.com/rss/markets",
- "https://www.livemint.com/rss/AI",
- "https://www.livemint.com/rss/insurance",
- "https://www.livemint.com/rss/budget",
- "https://www.livemint.com/rss/elections",
-]
-
-CNBC18_FEEDS = [
- "https://www.cnbctv18.com/commonfeeds/v1/cne/rss/latest.xml",
- "https://www.cnbctv18.com/commonfeeds/v1/cne/rss/india.xml",
- "https://www.cnbctv18.com/commonfeeds/v1/cne/rss/economy.xml",
- "https://www.cnbctv18.com/commonfeeds/v1/cne/rss/market.xml",
- "https://www.cnbctv18.com/commonfeeds/v1/cne/rss/business.xml",
- "https://www.cnbctv18.com/commonfeeds/v1/cne/rss/sports.xml",
- "https://www.cnbctv18.com/commonfeeds/v1/cne/rss/politics.xml",
- "https://www.cnbctv18.com/commonfeeds/v1/cne/rss/world.xml",
- "https://www.cnbctv18.com/commonfeeds/v1/cne/rss/education.xml",
- "https://www.cnbctv18.com/commonfeeds/v1/cne/rss/travel.xml",
- "https://www.cnbctv18.com/commonfeeds/v1/cne/rss/auto.xml",
- "https://www.cnbctv18.com/commonfeeds/v1/cne/rss/technology.xml",
- "https://www.cnbctv18.com/commonfeeds/v1/cne/rss/personal-finance.xml",
- "https://www.cnbctv18.com/commonfeeds/v1/cne/rss/web-stories.xml",
-]
-
-OTHER_FEEDS = [
- "http://www.business-standard.com/rss/todays-paper.rss",
- "https://www.forbesindia.com/commonfeeds/v1/frb/rss/blog.xml",
- "http://zeenews.india.com/rss/business.xml",
- "https://economictimes.indiatimes.com/rssfeedsdefault.cms",
- "https://news.google.com/rss?cf=all&hl=en-IN&topic=b&gl=IN&ceid=IN:en",
- "https://cfo.economictimes.indiatimes.com/rss/topstories",
- "https://cfo.economictimes.indiatimes.com/rss/recentstories",
- "https://cfo.economictimes.indiatimes.com/rss/corporate-finance",
- "https://cfo.economictimes.indiatimes.com/rss/esg",
- "https://cfo.economictimes.indiatimes.com/rss/cfo-tech",
- "https://cfo.economictimes.indiatimes.com/rss/governance-risk-compliance",
- "https://cfo.economictimes.indiatimes.com/rss/lateststories",
-]
-
-TE_SOURCES = {
- "te_stream": "https://tradingeconomics.com/stream",
- "te_economy": "https://tradingeconomics.com/stream?i=economy",
- "te_markets": "https://tradingeconomics.com/stream?i=markets",
- "te_india": "https://tradingeconomics.com/india/news",
-}
-
-ZERODHA_PULSE_URL = "https://pulse.zerodha.com/"
-
-
-# ==============================================================
-# TRADING ECONOMICS Selenium Scraper
-# ==============================================================
-
-JS_SCROLL_DOWN = "window.scrollTo(0, document.body.scrollHeight); return document.body.scrollHeight;"
-
-JS_EXTRACT_TE_NEWS = """
-var items = [];
-var listItems = document.querySelectorAll('li[id]');
-listItems.forEach(function(li) {
- var titleLink = li.querySelector('a[href]');
- if (!titleLink) return;
- var titleText = '';
- var bTag = titleLink.querySelector('b');
- if (bTag) { titleText = bTag.textContent.trim(); }
- else { titleText = titleLink.textContent.trim(); }
- if (!titleText) return;
- var url = titleLink.getAttribute('href') || '';
- if (url && !url.startsWith('http')) { url = 'https://tradingeconomics.com' + url; }
- var descEl = li.querySelector('.te-stream-item-description, span[style]');
- var description = descEl ? descEl.textContent.trim() : '';
- var dateEl = li.querySelector('small');
- var dateText = dateEl ? dateEl.textContent.trim() : '';
- var countryEl = li.querySelector('.te-stream-country');
- var country = countryEl ? countryEl.textContent.trim() : '';
- var categoryEl = li.querySelector('.te-stream-category');
- var category = categoryEl ? categoryEl.textContent.trim() : '';
- items.push({ title: titleText, description: description, date: dateText, url: url, country: country, category: category });
-});
-return JSON.stringify(items);
-"""
-
-
-class NewsClient:
- """
- Fetches news for a single ticker across a multi-year window by
- chunking into quarterly GDELT calls. Implements the same cache-first
- pattern as causal.test_causal_flow.test_causal_with_text.
-
- Usage:
- client = NewsClient("RELIANCE")
- articles = client.fetch("2022-01-01", "2026-04-30")
- # articles: List[dict] with keys: title, content, published, url, credibility_weight, source
- """
- USER_AGENT: str = (
- "Mozilla/5.0 (Windows NT 10.0; Win64; x64; rv:128.0) "
- "Gecko/20100101 Firefox/128.0"
- )
-
- def __init__(self, ticker: str):
- self.ticker = ticker.upper().replace(".NS", "")
- self.keyword = self._get_ticker_keyword(self.ticker)
-
- def _get_ticker_keyword(self, ticker: str) -> str:
- """Fetch a search-friendly name for the ticker from yfinance."""
- symbol = ticker if "." in ticker else f"{ticker}.NS"
- try:
- t = yf.Ticker(symbol)
- info = t.info
- name = info.get("shortName") or info.get("longName") or ticker
- # Clean up the name for better matching
- for suffix in [" Limited", " Ltd.", " Ltd", " Corp.", " Corp", " Inc.", " Inc"]:
- name = name.replace(suffix, "")
- return name.strip()
- except Exception as e:
- logger.warning(f"Failed to fetch yfinance info for {symbol}: {e}")
- return ticker
-
- def _normalise(self, raw: dict) -> dict:
- """Convert the raw dict format from news.py into a pipeline-ready article dict."""
- title = raw.get("title", "").strip()
- summary = raw.get("summary", "") or raw.get("description", "") or ""
- content = f"{title}. {summary}".strip(". ") if summary and summary != title else title
- pub_raw = raw.get("published", "") or raw.get("seendate", "")
- # Try to parse the published field
- pub_dt = None
- if pub_raw:
- try:
- from dateutil import parser as dp
- pub_dt = dp.parse(pub_raw)
- except Exception:
- pub_dt = None
- if pub_dt is None:
- pub_dt = datetime.datetime.now(datetime.timezone.utc)
- elif pub_dt.tzinfo is None:
- pub_dt = pub_dt.replace(tzinfo=datetime.timezone.utc)
-
- return {
- "title": title,
- "content": content,
- "published": pub_dt.isoformat(),
- "url": raw.get("link", "") or raw.get("url", ""),
- }
-
- def fetch(
- self,
- start: str,
- end: str,
- include_pulse: bool = False,
- include_te: bool = False,
- ) -> List[dict]:
- """
- Fetch all news for self.ticker between start and end.
- Aggregates RSS, Pulse, and (optionally) Trading Economics.
- """
-
- logger.info(
- f"Fetching news for {self.ticker} ('{self.keyword}') "
- f"from {start} to {end}..."
- )
-
- all_raw: List[dict] = []
-
- # 1. RSS Feeds
- logger.info("Searching RSS feeds...")
- rss_articles = (
- self.search_rss_feeds(LIVEMINT_FEEDS, self.keyword, "LiveMint")
- + self.search_rss_feeds(CNBC18_FEEDS, self.keyword, "CNBC18")
- + self.search_rss_feeds(OTHER_FEEDS, self.keyword, "OtherFeeds")
- )
- all_raw.extend(rss_articles)
-
- # 2. Zerodha Pulse
- if include_pulse:
- logger.info("Fetching Zerodha Pulse...")
- pulse_news = self.scrape_pulse()
- all_raw.extend(pulse_news)
-
- # 3. Trading Economics (Optional)
- if include_te:
- logger.info("Fetching Trading Economics...")
- te_data = self.scrape_all_te(headless=True, scroll_count=1)
- for items_list in te_data.values():
- all_raw.extend(items_list)
-
- # # Filter by date window
- # start_dt = datetime.datetime(
- # *[int(x) for x in start.split("-")], tzinfo=datetime.timezone.utc
- # )
- # end_dt = datetime.datetime(
- # *[int(x) for x in end.split("-")], tzinfo=datetime.timezone.utc
- # )
- # in_window = self._filter_by_window(all_raw, start_dt, end_dt)
- # logger.info(f"Collected {len(in_window)} articles in window.")
-
- # # Normalise and deduplicate by title
- # seen_titles: set = set()
- # articles: List[dict] = []
- # for raw in all_raw:
- # title = raw.get("title", "").strip().lower()
- # if not title or title in seen_titles:
- # continue
- # seen_titles.add(title)
- # articles.append(self._normalise(raw))
-
- # Sort chronologically
- all_raw.sort(key=lambda a: str(a.get("published") or a.get("date") or ""))
-
- logger.info(
- f"Final corpus: {len(all_raw)} unique articles for {self.ticker}."
- )
-
- return all_raw
-
- def search_rss_feeds(self, feeds: list, search_keyword: str, feed_name: str) -> list:
- """Search for a keyword across a list of RSS feed URLs."""
- print(f"[DEBUG] search_rss_feeds - Searching {feed_name} for keyword: '{search_keyword}'")
- results = []
- for feed_url in feeds:
- print(f"[DEBUG] search_rss_feeds - Parsing URL: {feed_url}")
- try:
- feed = feedparser.parse(feed_url)
- if not feed.entries:
- print(f"[DEBUG] search_rss_feeds - No entries found for {feed_url}")
- for entry in feed.entries:
- title = entry.get("title", "")
- summary = entry.get("summary", "")
- if not title:
- continue
- if search_keyword.lower() in title.lower() or search_keyword.lower() in summary.lower():
- print(f"[DEBUG] search_rss_feeds - Match found: {title[:60]}...")
- results.append({
- "feed_name": feed_name,
- "title": title,
- "link": entry.get("link", ""),
- "published": entry.get("published", ""),
- "summary": summary,
- })
- except Exception as e:
- print(f"[news.py] RSS error ({feed_url}): {e}")
- print(f"[DEBUG] search_rss_feeds - {feed_name} done. Found {len(results)} items.")
- return results
-
- def scrape_all_te(self, headless: bool = True, scroll_count: int = 3) -> dict:
- """Scrape all Trading Economics news sources. Returns dict of source_key -> list of items."""
- driver = self._create_driver(headless=headless)
- results = {}
- scraped_at = datetime.datetime.now(datetime.timezone.utc).isoformat()
- try:
- for key, url in TE_SOURCES.items():
- try:
- items = self._scrape_te_page(driver, url, scroll_count=scroll_count)
- for item in items:
- item["source"] = key
- item["scraped_at"] = scraped_at
- results[key] = items
- except Exception as e:
- print(f" [TE] Error scraping {key}: {e}")
- results[key] = []
- try:
- driver.quit()
- except Exception:
- pass
- driver = self._create_driver(headless=headless)
- finally:
- try:
- driver.quit()
- except Exception:
- pass
- return results
-
- def scrape_pulse(self) -> list:
- """Scrape latest news from Zerodha Pulse using requests + BeautifulSoup."""
- print("[news.py] Scraping Zerodha Pulse...")
- headers = {
- "User-Agent": self.USER_AGENT,
- "Accept": "text/html,application/xhtml+xml,*/*;q=0.8",
- "Accept-Language": "en-US,en;q=0.5",
- }
- try:
- resp = requests.get(ZERODHA_PULSE_URL, headers=headers, timeout=30)
- resp.raise_for_status()
- except Exception as e:
- print(f"[news.py] Zerodha Pulse error: {e}")
- return []
-
- soup = BeautifulSoup(resp.text, "html.parser")
- items = []
- scraped_at = datetime.datetime.now(datetime.timezone.utc).isoformat()
-
- for li in soup.select("li.item"):
- title_el = li.select_one("h2.title a")
- if not title_el:
- continue
- title = title_el.get_text(strip=True)
- url = title_el.get("href", "")
- desc_el = li.select_one("div.desc")
- description = desc_el.get_text(strip=True) if desc_el else ""
- date_el = li.select_one("span.date")
- date_text = date_el.get_text(strip=True) if date_el else ""
- feed_el = li.select_one("span.feed")
- feed = feed_el.get_text(strip=True) if feed_el else ""
- items.append({
- "title": title,
- "description": description,
- "published": date_text,
- "link": url,
- "summary": description,
- "publisher": feed,
- "scraped_at": scraped_at,
- "source": "pulse"
- })
-
- print(f"[news.py] Zerodha Pulse: extracted {len(items)} items.")
- return items
-
- def _filter_by_window(
- self,
- articles: list,
- from_dt: datetime.datetime,
- to_dt: datetime.datetime,
- ) -> list:
- """
- Drop articles whose parsed `published` timestamp falls outside [from_dt, to_dt].
- Articles with unparseable or missing dates are kept (conservative).
- """
- from dateutil import parser as dp
-
- def _to_utc(dt):
- if dt.tzinfo is None:
- return dt.replace(tzinfo=datetime.timezone.utc)
- return dt.astimezone(datetime.timezone.utc)
-
- from_utc = _to_utc(from_dt)
- to_utc = _to_utc(to_dt)
-
- filtered = []
- for art in articles:
- pub = art.get("published", "")
- if not pub:
- filtered.append(art)
- continue
- try:
- dt = _to_utc(dp.parse(pub))
- if from_utc <= dt <= to_utc:
- filtered.append(art)
- except Exception:
- filtered.append(art) # keep on parse failure
- return filtered
-
- def _create_driver(self, headless: bool = True):
- """Create a headless Firefox webdriver."""
- options = FirefoxOptions()
- if headless:
- options.add_argument("--headless")
- options.set_preference("general.useragent.override", self.USER_AGENT)
- options.set_preference("dom.webdriver.enabled", False)
- options.set_preference("useAutomationExtension", False)
- service = FirefoxService(log_output=os.devnull)
- driver = webdriver.Firefox(options=options, service=service)
- driver.set_page_load_timeout(60)
- return driver
-
- def _scrape_te_page(self, driver, url: str, scroll_count: int = 3, scroll_pause: float = 2.0) -> list:
- """Scrape a single Trading Economics news page."""
- print(f" [TE] Loading: {url}")
- driver.get(url)
- WebDriverWait(driver, 20).until(EC.presence_of_element_located((By.TAG_NAME, "body")))
- time.sleep(3)
- prev_height = 0
- for i in range(scroll_count):
- new_height = driver.execute_script(JS_SCROLL_DOWN)
- if new_height == prev_height:
- break
- prev_height = new_height
- time.sleep(scroll_pause)
- raw = driver.execute_script(JS_EXTRACT_TE_NEWS)
- items = json.loads(raw)
- print(f" [TE] Extracted {len(items)} items from {url}")
- return items
-
-
-if __name__ == "__main__":
- """Simple test runner for NewsClient."""
- logging.basicConfig(
- level=logging.INFO,
- format="%(asctime)s [%(levelname)s] %(name)s: %(message)s"
- )
-
- ticker = "RELIANCE"
- client = NewsClient(ticker)
-
- # Test Fetch (Aggregates RSS and Pulse)
- print(f"\n--- Testing Fetch for {ticker} ---")
- start_date = (datetime.datetime.now() - datetime.timedelta(days=7)).strftime("%Y-%m-%d")
- end_date = datetime.datetime.now().strftime("%Y-%m-%d")
-
- # We include_te=False by default as it requires Selenium/Firefox
- articles = client.fetch(start_date, end_date)
- print(f"Fetched {len(articles)} articles.")
-
- if articles:
- # Save results to debug_data/
- with open("news.json", "w", encoding="utf-8") as f:
- json.dump(articles, f, indent=4, ensure_ascii=False)
- for a in articles[:3]:
- print(f" - [{a['published']}] {a['title']}")
- else:
- print("No articles found in the specified window.")
diff --git a/singular_ticker_causal/data_sources/nseconnect/__init__.py b/singular_ticker_causal/data_sources/nseconnect/__init__.py
deleted file mode 100644
index c6f424bece3ca95dcf05a69a47d006e553b2640c..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/data_sources/nseconnect/__init__.py
+++ /dev/null
@@ -1,25 +0,0 @@
-"""
- The MIT License (MIT)
-
- Copyright (c) 2014 Noufal Nazar
-
- Permission is hereby granted, free of charge, to any person obtaining a copy
- of this software and associated documentation files (the "Software"), to deal
- in the Software without restriction, including without limitation the rights
- to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
- copies of the Software, and to permit persons to whom the Software is
- furnished to do so, subject to the following conditions:
-
- The above copyright notice and this permission notice shall be included in all
- copies or substantial portions of the Software.
-
- THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
- IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
- FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
- AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
- LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
- OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
- SOFTWARE.
-"""
-__VERSION__='2.0.1'
-from .nse import Nse
diff --git a/singular_ticker_causal/data_sources/nseconnect/bases.py b/singular_ticker_causal/data_sources/nseconnect/bases.py
deleted file mode 100644
index e87e10ec5c16acfc025886c608a22265a23e5d89..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/data_sources/nseconnect/bases.py
+++ /dev/null
@@ -1,72 +0,0 @@
-"""
- The MIT License (MIT)
-
- Copyright (c) 2014 Noufal Nazar
-
- Permission is hereby granted, free of charge, to any person obtaining a copy
- of this software and associated documentation files (the "Software"), to deal
- in the Software without restriction, including without limitation the rights
- to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
- copies of the Software, and to permit persons to whom the Software is
- furnished to do so, subject to the following conditions:
-
- The above copyright notice and this permission notice shall be included in all
- copies or substantial portions of the Software.
-
- THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
- IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
- FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
- AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
- LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
- OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
- SOFTWARE.
-"""
-
-from abc import ABCMeta, abstractmethod
-import six
-
-
-class AbstractBaseExchange(six.with_metaclass(ABCMeta, object)):
-
- @abstractmethod
- def get_stock_codes(self):
- """
- :return: list of tuples with stock code and stock name
- """
- raise NotImplementedError
-
- @abstractmethod
- def is_valid_code(self, code):
- """
- :return: True, if it is a valid stock code, else False
- """
- raise NotImplementedError
-
- @abstractmethod
- def get_quote(self, code):
- """
- :param code: a stock code
- :return: a dictionary which contain detailed stock code.
- """
- raise NotImplementedError
-
- @abstractmethod
- def get_top_gainers(self):
- """
- :return: a sorted list of codes of top gainers
- """
- raise NotImplementedError
-
- @abstractmethod
- def get_top_losers(self):
- """
- :return: a sorted list of codes of top losers
- """
- raise NotImplementedError
-
- @abstractmethod
- def __str__(self):
- """
- :return: market name
- """
- raise NotImplementedError
diff --git a/singular_ticker_causal/data_sources/nseconnect/cleaners.py b/singular_ticker_causal/data_sources/nseconnect/cleaners.py
deleted file mode 100644
index ec382b6e40419975769106fcbaf111afa2a24733..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/data_sources/nseconnect/cleaners.py
+++ /dev/null
@@ -1,51 +0,0 @@
-"""
-Module for various data structure cleaning tasks
-"""
-from datetime import datetime
-dirty_data = """
-{
- "fname": "Jon",
- "lname": "Doe",
- "age": 20,
- "str_age": "20",
- "pi": 3.1415927,
- "str_pi": "3.1415927",
- "dob": "01-Jan-2023",
- "mobile": [
- {
- "id": "Home",
- "number": "123456789"
- },
- {
- "id": "office",
- "number": "987645321"
- }
- ]
-}
-"""
-
-def parse_values(obj):
- for key, value in obj.items():
- if isinstance(value, str):
- # Try to parse as datetime if the string matches the format
- date_formats = ["%d-%b-%Y", "%d-%m-%Y", "%Y-%m-%d %H:%M:%S", "%Y-%m-%d"]
- for date_format in date_formats:
- try:
- obj[key] = datetime.strptime(value, date_format)
- break
- except ValueError:
- pass
- else:
- # If the string couldn't be parsed as datetime, try numeric conversion
- try:
- obj[key] = int(value)
- except ValueError:
- try:
- obj[key] = float(value)
- except ValueError:
- pass
- elif isinstance(value, dict):
- obj[key] = parse_values(value)
- elif isinstance(value, list):
- obj[key] = [parse_values(item) if isinstance(item, dict) else item for item in value]
- return obj
diff --git a/singular_ticker_causal/data_sources/nseconnect/datemgr.py b/singular_ticker_causal/data_sources/nseconnect/datemgr.py
deleted file mode 100644
index a698d314f5b8381e1545cc897f4c07d14afbfcc8..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/data_sources/nseconnect/datemgr.py
+++ /dev/null
@@ -1,104 +0,0 @@
-import datetime as dt
-from dateutil.relativedelta import relativedelta
-from dateutil.parser import parse
-from dateutil import rrule
-from .errors import DateFormatError
-
-
-def get_nearest_business_day(d):
- """ takes datetime object"""
- if d.isoweekday() == 7 or d.isoweekday() == 6:
- d = d - relativedelta(days=1)
- return get_nearest_business_day(d)
-
- # republic day
- elif d.month == 1 and d.day == 26:
- d = d - relativedelta(days=1)
- return get_nearest_business_day(d)
- # labour day
- elif d.month == 5 and d.day == 1:
- d = d - relativedelta(days=1)
- return get_nearest_business_day(d)
- # independece day
- elif d.month == 8 and d.day == 15:
- d = d - relativedelta(days=1)
- return get_nearest_business_day(d)
- # Gandhi Jayanti
- elif d.month == 10 and d.day == 2:
- d = d - relativedelta(days=1)
- return get_nearest_business_day(d)
- # chirstmas
- elif d.month == 12 and d.day == 25:
- d = d - relativedelta(days=1)
- return get_nearest_business_day(d)
- else:
- return d
-
-def is_known_holiday(d):
- """accepts datetime/date object and returns boolean"""
- if type(d) == dt.datetime:
- d = d.date()
- elif type(d) != dt.date:
- raise DateFormatError("only date objects or datetime objects")
- else:
- # fine do nothing
- pass
-
- # declare the list of holidays here.
- # republic day.
- if d.month == 1 and d.day == 26:
- return True
- # labour day
- elif d.month == 5 and d.day == 1:
- d = d - relativedelta(days=1)
- return get_nearest_business_day(d)
- # independence day
- elif d.month == 8 and d.day == 15:
- return True
- # gandhi jayanti
- elif d.month == 10 and d.day == 2:
- return True
- # christmas
- elif d.month == 12 and d.day == 25:
- return True
- else:
- return False
-
-def mkdate(d):
- """tries its best to return a valid date. it can accept pharse like today,
- yesterday, day before yesterday etc.
- """
- # check if the it == a string
- return_date = ""
- if type(d) is str:
- if d == "today":
- return_date = dt.date.today()
- elif d == "yesterday":
- return_date = dt.date.today() - relativedelta(days=1)
- elif d == "day before yesterday":
- return_date = dt.date.today() - relativedelta(days=2)
- else:
- return_date = parse(d, dayfirst=True).date()
- elif type(d) == dt.datetime:
- return_date = d.date()
- elif type(d) == dt.date:
- return d
- else:
- raise DateFormatError("wrong date format %s" % str(d))
- # check if future date.
- return return_date
-
-def usable_date(d):
- """accepts fuzzy format and returns most sensible date"""
- return get_nearest_business_day(mkdate(d))
-
-def get_date_range(frm, to, skip_dates=[]):
- """accepts fuzzy format date and returns business adjusted date ranges"""
- # for x in rrule.rrule(rrule.DAILY, dtstart=s, until=dt.datetime.now(), byweekday=[0, 1, 2, 3, 4]): print(x)
- frm = usable_date(frm)
- to = usable_date(to)
- datelist = []
- for date in rrule.rrule(rrule.DAILY, dtstart=frm, until=to, byweekday=[0, 1, 2, 3, 4]):
- if not is_known_holiday(date):
- datelist.append(date.date())
- return datelist
diff --git a/singular_ticker_causal/data_sources/nseconnect/downloader.py b/singular_ticker_causal/data_sources/nseconnect/downloader.py
deleted file mode 100644
index 9a2dc07bf77e4e7c4cf142acac3030301cf36d36..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/data_sources/nseconnect/downloader.py
+++ /dev/null
@@ -1,116 +0,0 @@
-"""
- The MIT License (MIT)
-
- Copyright (c) 2014 Noufal Nazar
-
- Permission is hereby granted, free of charge, to any person obtaining a copy
- of this software and associated documentation files (the "Software"), to deal
- in the Software without restriction, including without limitation the rights
- to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
- copies of the Software, and to permit persons to whom the Software is
- furnished to do so, subject to the following conditions:
-
- The above copyright notice and this permission notice shall be included in all
- copies or substantial portions of the Software.
-
- THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
- IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
- FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
- AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
- LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
- OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
- SOFTWARE.
-
-"""
-import io
-import os
-import zipfile
-import datetime as dt
-from urllib.request import Request
-from .datemgr import mkdate, get_date_range
-from .nse import Nse
-from abc import ABCMeta, abstractmethod
-
-class BaseBhavcopyDownloader(metaclass=ABCMeta):
- """Base class for all types of bhavcopy downloader"""
- def __init__(self, from_date, to_date=dt.datetime.now().date(), skip_dates=[]):
- """accepts date in fuzzy format"""
- self.bhavcopy_base_url = "https://www.nseindia.com/content/historical/EQUITIES/%s/%s/cm%s%s%sbhav.csv.zip"
- self.bhavcopy_base_filename = "cm%s%s%sbhav.csv"
- self.from_date = from_date
- self.to_date = to_date
- self.skip_dates = skip_dates
- self.nse = Nse()
- self.dates = self.generate_dates()
-
- def generate_dates(self):
- return get_date_range(self.from_date, self.to_date, skip_dates=self.skip_dates)
-
- def get_bhavcopy_url(self, d):
- """accept date and return bhavcopy url"""
- day_of_month = d.strftime("%d")
- mon = d.strftime("%b").upper()
- year = d.year
- url = self.bhavcopy_base_url % (year, mon, day_of_month, mon, year)
- return url
-
- def get_bhavcopy_filename(self, d):
- """for a given date generate bhavcopy filename"""
- day_of_month = d.strftime("%d")
- mon = d.strftime("%b").upper()
- year = d.year
- filename = self.bhavcopy_base_filename % (day_of_month, mon, year)
- return filename
-
- def download_one(self, d):
- """download bhavcopy for the given date"""
- # this will keep this method usable for any arbitrary date.
- d = mkdate(d)
- # ex_url = "https://www.nseindia.com/content/historical/EQUITIES/2011/NOV/cm08NOV2011bhav.csv.zip"
- url = self.get_bhavcopy_url(d)
- print(url)
- filename = self.get_bhavcopy_filename(d)
- # response = requests.get(url, headers=self.headers)
- response = self.nse.opener.open(Request(url, None, self.nse.headers))
- zip_file_handle = io.BytesIO(response.read())
- zf = zipfile.ZipFile(zip_file_handle)
- return zf.read(filename).decode("utf-8")
-
- @abstractmethod
- def download(self):
- pass
-
- @abstractmethod
- def update(self):
- pass
-
-
-class BhavcopyFileSystemDownloader(BaseBhavcopyDownloader):
- def __init__(self, directory, *args, **kwargs):
- if (os.path.exists(directory) and os.path.isdir(directory) and os.access(directory, os.W_OK)):
- super().__init__(*args, **kwargs)
- self.directory = directory
- else:
- raise Exception("directory path must be valid and writtable, please check manually")
-
- def download(self):
- for date in self.dates:
- print("downloading for " + str(date))
- try:
- content = self.download_one(date)
- except Exception as err:
- print("unable to download for the date: %s" % date.strftime("%Y-%m-%d"))
- else:
- fh = open(self.directory + "/" + date.strftime("%Y-%m-%d") + ".csv", "w")
- fh.write(content)
- fh.close()
-
- def update(self):
- pass
-
-
-if __name__ == '__main__':
- b = BhavcopyFileSystemDownloader(directory="/tmp/bhavcopy", from_date="01-01-2018")
- b.download()
-
-# https://stackoverflow.com/questions/49183801/ssl-certificate-verify-failed-with-urllib
\ No newline at end of file
diff --git a/singular_ticker_causal/data_sources/nseconnect/errors.py b/singular_ticker_causal/data_sources/nseconnect/errors.py
deleted file mode 100644
index 027a54050382763559edbe8da4546c6ce5b14f4a..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/data_sources/nseconnect/errors.py
+++ /dev/null
@@ -1,8 +0,0 @@
-class BhavcopyNotAvailableError(Exception):
- """this error could occur in case you download bhavcopy for the dates
- when the market was close"""
- pass
-
-class DateFormatError(Exception):
- """in case the date format is errorneous"""
- pass
\ No newline at end of file
diff --git a/singular_ticker_causal/data_sources/nseconnect/nse.py b/singular_ticker_causal/data_sources/nseconnect/nse.py
deleted file mode 100644
index c111f21e3414ceecba90fa5fe9f7e95661a8640f..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/data_sources/nseconnect/nse.py
+++ /dev/null
@@ -1,624 +0,0 @@
-"""
- The MIT License (MIT)
-
- Copyright (c) 2014 Noufal Nazar
-
- Permission is hereby granted, free of charge, to any person obtaining a copy
- of this software and associated documentation files (the "Software"), to deal
- in the Software without restriction, including without limitation the rights
- to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
- copies of the Software, and to permit persons to whom the Software is
- furnished to do so, subject to the following conditions:
-
- The above copyright notice and this permission notice shall be included in all
- copies or substantial portions of the Software.
-
- THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
- IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
- FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
- AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
- LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
- OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
- SOFTWARE.
-
-"""
-
-import csv
-from .bases import AbstractBaseExchange
-from .urls import (
- STOCKS_CSV_URL, STOCKS_IN_INDEX_URL, QUOTE_API_URL, QUOTE_DRIVATIVE_URL,
- TOP_GAINERS_URL, TOP_LOSERS_URL, ALL_INDICES_URL,
- FIFTYTWO_WEEK_HIGH_URL, FIFTYTWO_WEEK_LOW_URL
-)
-from .ua import Session
-from .utils import cast_intfloat_string_values_to_intfloat
-
-class Nse(AbstractBaseExchange):
- """
- class which implements all the functionality for
- National Stock Exchange
- """
- __CODECACHE__ = None
-
- def __init__(self, session_refresh_interval=120):
- """Initialize a new NSE object.
- Initializes a session management for making API calls to NSE (National Stock Exchange).
- Args:
- session_refresh_interval (int, optional): Time interval in seconds after which the session
- should be refreshed. Defaults to 120 seconds.
- Note:
- The session refresh interval helps maintain an active connection with NSE servers by
- periodically creating a new session to prevent timeouts.
- """
-
- self.session_refresh_interval = session_refresh_interval
- self.session = Session(session_refresh_interval)
-
- #############################
- ### STOCKS APIS ###
- #############################
-
- def get_stock_codes(self):
- """Gets a list of stock codes traded in NSE.
-
- This function fetches stock data from NSE's CSV endpoint and extracts the stock symbols.
-
- Returns:
- list: A list of strings containing stock symbols traded on NSE.
-
- Example:
- >>> nse = Nse()
- >>> codes = nse.get_stock_codes()
- >>> print(codes[:5])
- ['20MICRONS', '3IINFOTECH', '3MINDIA', '3PLAND', '63MOONS']
- """
- res = self.session.fetch(STOCKS_CSV_URL)
- csv_content = res.text.splitlines()
- symbols = []
- csv_reader = csv.DictReader(csv_content)
- for row in csv_reader:
- symbols.append(row['SYMBOL'])
- return symbols
-
- def is_valid_code(self, code):
- """Checks if a given stock code is valid.
-
- This method validates whether the provided stock code exists in the list of valid
- stock codes from NSE (National Stock Exchange).
-
- Args:
- code (str): Stock code/symbol to validate.
-
- Returns:
- bool: True if the code is valid, False otherwise.
-
- Example:
- >>> nse = NSE()
- >>> nse.is_valid_code("INFY")
- True
- >>> nse.is_valid_code("INVALID")
- False
- """
- stock_codes = self.get_stock_codes()
- return code.upper() in stock_codes
-
- def get_quote(self, code, all_data=False):
- """Gets the stock quote for a given NSE stock symbol.
-
- This function fetches real-time or delayed quote data from NSE for the specified stock code.
-
- Args:
- code (str): NSE stock symbol/code for which quote is to be fetched
- all_data (bool, optional): If True returns complete quote data, if False returns only price info.
- Defaults to False.
-
- Returns:
- dict: A dictionary containing quote data.
-
- Raises:
- requests.exceptions.RequestException: If there is an error in HTTP request
- ValueError: If the response JSON is invalid
-
- Example:
- >>> nse = Nse()
- >>> nse.get_quote('abb')
- {
- 'lastPrice': 5189.1,
- 'change': 70.55,
- 'pChange': 1.38,
- 'previousClose': 5118.55,
- 'open': 5160,
- 'close': 5187.65,
- 'vwap': 5162.91,
- 'stockIndClosePrice': 0,
- 'lowerCP': 4606.7,
- 'upperCP': 5630.4,
- 'pPriceBand': 'No Band',
- 'basePrice': 5118.55,
- 'intraDayHighLow': {'min': 5101, 'max': 5218.45, 'value': 5189.1},
- 'weekHighLow': {'min': 4890}
- }
- """
- code = code.upper()
- # TODO: implement if the code is valid
- res = self.session.fetch(QUOTE_API_URL % code)
- res = res.json()['priceInfo'] if all_data is False else res.json()
- return cast_intfloat_string_values_to_intfloat(res)
-
- def get_52_week_high(self):
- """Retrieves a list of stocks that have hit their 52-week high.
-
- This method fetches data for stocks that have reached new 52-week high prices on the NSE.
-
- Returns:
- list[dict]: A list of dictionaries containing 52-week high data.
-
- Example:
- >>> nse.get_52_week_high()
- [{'symbol': 'AVANTIFEED',
- 'series': 'EQ',
- 'comapnyName': 'Avanti Feeds Limited',
- 'new52WHL': 899,
- 'prev52WHL': 849.9,
- 'prevHLDate': '13-Mar-2025',
- 'ltp': 887,
- 'prevClose': 842.55,
- 'change': 44.45,
- 'pChange': 5.28},
- {...}
- ]
- """
- res = self.session.fetch(FIFTYTWO_WEEK_HIGH_URL)
- json_response = res.json()
- # Handle the new API response structure which has dataLtpGreater20 and dataLtpLess20 fields
- data = cast_intfloat_string_values_to_intfloat(json_response)
-
- # Check if the old structure with 'data' key exists
- if 'data' in data:
- return data['data']
-
- # Otherwise, extract and combine the lists from the new structure
- result = []
- if 'dataLtpGreater20' in data:
- result.extend(data['dataLtpGreater20'])
- if 'dataLtpLess20' in data:
- result.extend(data['dataLtpLess20'])
- return result
-
- def get_52_week_low(self):
- """Retrieves a list of stocks that have hit their 52-week low.
-
- This method fetches data for stocks that have reached new 52-week low prices on the NSE.
-
- Returns:
- list[dict]: A list of dictionaries containing 52-week low data.
-
- Example:
- >>> nse.get_52_week_low()
- [{'symbol': 'AVANTIFEED',
- 'series': 'EQ',
- 'comapnyName': 'Avanti Feeds Limited',
- 'new52WHL': 899,
- 'prev52WHL': 849.9,
- 'prevHLDate': '13-Mar-2025',
- 'ltp': 887,
- 'prevClose': 842.55,
- 'change': 44.45,
- 'pChange': 5.28},
- {...}
- ]
- """
- res = self.session.fetch(FIFTYTWO_WEEK_LOW_URL)
- json_response = res.json()
- # Handle the new API response structure which has dataLtpGreater20 and dataLtpLess20 fields
- data = cast_intfloat_string_values_to_intfloat(json_response)
-
- # Check if the old structure with 'data' key exists
- if 'data' in data:
- return data['data']
-
- # Otherwise, extract and combine the lists from the new structure
- result = []
- if 'dataLtpGreater20' in data:
- result.extend(data['dataLtpGreater20'])
- if 'dataLtpLess20' in data:
- result.extend(data['dataLtpLess20'])
- return result
-
- #############################
- ### INDEX APIS ###
- #############################
-
- def get_index_quote(self, index="NIFTY 50"):
- """Gets the quote for a specific index from NSE.
-
- This function retrieves detailed quote information for a given index code from the
- National Stock Exchange (NSE) of India.
-
- Args:
- index (str): The index code/symbol (e.g. "NIFTY 50", "BANKNIFTY", etc.)
-
- Returns:
- dict: A dictionary containing index quote details
-
- Raises:
- Exception: If the provided index code is invalid or not found
-
- Example:
- >>> nse = NSE()
- >>> nse.get_index_quote("NIFTY 50")
- {
- 'key': 'BROAD MARKET INDICES',
- 'index': 'NIFTY 50',
- 'last': 22508.75,
- 'variation': 111.55,
- 'percentChange': 0.5,
- 'open': 22353.15,
- 'high': 22577.0,
- 'low': 22353.15,
- 'previousClose': 22397.2,
- 'yearHigh': 26277.35,
- 'yearLow': 21281.45,
- # ... additional fields omitted for brevity
- }
- """
-
- url = ALL_INDICES_URL
- all_index_quote = self.get_all_index_quote()
- index_list = [ i['indexSymbol'] for i in all_index_quote]
- index = index.upper()
- index = ' '.join(index.split())
- if index in index_list:
- response = list(filter(lambda idx: idx['indexSymbol'] == index, all_index_quote))[0]
- return cast_intfloat_string_values_to_intfloat(response)
- else:
- raise Exception('Wrong index code')
-
- def get_index_list(self):
- """Gets a list of all NSE index symbols.
-
- This method fetches all available NSE (National Stock Exchange) index symbols by
- extracting the 'indexSymbol' from the complete index quote data.
-
- Returns:
- list: A list of strings containing index symbols (e.g., ['NIFTY 50', 'NIFTY BANK', ...])
-
- Examples:
- >>> nse = Nse()
- >>> indices = nse.get_index_list()
- >>> print(indices)
- ['NIFTY 50', 'NIFTY BANK', 'NIFTY IT', ...]
- """
- return [ i['indexSymbol'] for i in self.get_all_index_quote()]
-
- def get_all_index_quote(self):
- """Gets information for all NSE indices in one request.
-
- This method fetches quotes and information for all available indices on the
- National Stock Exchange (NSE) through a single API call.
-
- Returns:
- list[dict]: A list of dictionaries where each dictionary contains quote
- information for an index. The quote information includes details like
- index name, current value, change, percentage change etc.
-
- Example:
- >>> nse = Nse()
- >>> quotes = nse.get_all_index_quote()
- >>> quotes # Sample output
- [
- {
- 'key': 'BROAD MARKET INDICES',
- 'index': 'NIFTY 50',
- 'indexSymbol': 'NIFTY 50',
- 'last': 22508.75,
- 'variation': 111.55,
- 'percentChange': 0.5,
- 'open': 22353.15,
- ...
- },
- # ... additional indices follow
- ]
-
- Raises:
- URLError: If there is an error accessing the NSE API endpoint
- ValueError: If the response JSON cannot be parsed properly
- """
- url = ALL_INDICES_URL
- res = self.session.fetch(url)
- return res.json()['data']
-
- def get_top_gainers(self, index="NIFTY"):
- """Gets the list of top gaining stocks for the specified index.
-
- This function retrieves real-time data for stocks that have gained the most value
- during the current trading day. It can filter results by different indices.
-
- Args:
- index (str, optional): The index to get top gainers for. Defaults to "NIFTY".
- Valid values are:
- - NIFTY: Nifty 50 index
- - BANKNIFTY: Bank Nifty index
- - NIFTYNEXT50: Nifty Next 50 index
- - SecGtr20: Securities greater than 20
- - SecLwr20: Securities lower than 20
- - FNO: Futures & Options
- - ALL: All stocks
-
- Returns:
- list[dict]: List of dictionaries containing top gainer details.
-
- Raises:
- ConnectionError: If unable to fetch data from NSE
-
- Example:
- >>> nse = Nse()
- >>> gainers = nse.get_top_gainers()
- >>> gainers[0] # Sample output
- {
- 'symbol': 'DRREDDY',
- 'series': 'EQ',
- 'open_price': 1107.9,
- 'high_price': 1154.1,
- 'low_price': 1101.5,
- 'ltp': 1151.5,
- 'prev_price': 1107.95,
- 'net_price': 3.93,
- 'trade_quantity': 2714559,
- 'turnover': 31016.01,
- 'market_type': 'N',
- 'ca_ex_dt': '28-Oct-2024',
- 'ca_purpose': 'Face Value Split (Sub-Division) - From Rs 5/- Per Share To Re 1/- Per Share',
- 'perChange': 3.93
- }
- """
- return self._get_top_gainers_losers('gainers', index)
-
- def get_top_losers(self, index="NIFTY"): # Changed from None to "NIFTY"
- """Gets the top losers from specified index from NSE.
-
- The function fetches real-time data for stocks that have declined the most in terms
- of percentage change compared to their previous closing price.
-
- Args:
- index (str, optional): Index name for which top losers are to be fetched.
- Available options:
- - NIFTY (Default)
- - BANKNIFTY
- - NIFTYNEXT50
- - SecGtr20
- - SecLwr20
- - FNO
- - ALL
-
- Returns:
- list: List of dictionaries containing stock information with following keys:
-
- Raises:
- URLError: When unable to connect to NSE
- ValueError: When invalid index is provided
-
- Examples:
- >>> from nseconnect import Nse
- >>> nse = Nse()
- >>> losers = nse.get_top_losers()
- >>> losers[0]
- {'symbol': 'TATAMOTORS', 'series': 'EQ', 'openPrice': 375.0, ...}
- """
- return self._get_top_gainers_losers('losers', index) # Changed from 'gainers' to 'losers'
-
- def get_advances_declines(self, index='nifty 50'):
- """Gets the advances/declines data for given index.
- This method provides the number of stocks advancing and declining in a given index
- on NSE at any given point of time.
- Args:
- index (str, optional): Name of the index. Defaults to 'nifty 50'.
- Valid values include 'NIFTY 50', 'NIFTY BANK', etc.
- Returns:
- dict: A dictionary with two keys:
- - 'advances': Number of advancing stocks in the index
- - 'declines': Number of declining stocks in the index
- Examples:
- >>> nse = Nse()
- >>> nse.get_advances_declines(index="NIFTY BANK")
- {'advances': 7, 'declines': 4}
- Note:
- The method is case-insensitive for the index parameter.
- """
-
- # fixing this
- index = index.upper()
- index_quote = self.get_index_quote(index)
- return {'advances': index_quote['advances'], 'declines': index_quote['declines']}
-
- def get_stocks_in_index(self, index="NIFTY 50"):
- """Gets the list of symbols of stocks included in the specified NSE index.
- The function retrieves the current constituents of a given NSE index like NIFTY 50,
- NIFTY BANK etc. and returns their stock symbols.
- Args:
- index (str, optional): Name of the NSE index. Defaults to "NIFTY 50".
- Possible values: "NIFTY 50", "NIFTY BANK", "NIFTY IT" etc.
- Returns:
- list: List of stock symbols (str) that are part of the specified index.
- Raises:
- URLError: If unable to connect to NSE server
- ValueError: If invalid index name is provided
- Examples:
- >>> nse = Nse()
- >>> nse.get_stocks_in_index("NIFTY 50")
- ['ADANIPORTS', 'ASIANPAINT', 'AXISBANK', 'BAJAJ-AUTO', 'BAJFINANCE', ...]
- >>> nse.get_stocks_in_index("NIFTY BANK")
- ['AUBANK', 'AXISBANK', 'BANDHANBNK', 'FEDERALBNK', 'HDFCBANK', ...]
- """
-
- index = index.upper()
- url = STOCKS_IN_INDEX_URL % index
- res = self.session.fetch(url)
- res_dict = res.json()
- return [stock['symbol'] for stock in res_dict['data']][1:]
-
- def get_stock_quote_in_index(self, index="NIFTY 50", include_index=False):
- """Gets stock quotes for all stocks in a given index.
- This function fetches real-time quotes for all stocks that are part of the specified index
- from NSE (National Stock Exchange).
- Args:
- index (str, optional): The name of the index. Defaults to "NIFTY 50".
- include_index (bool, optional): Whether to include the index itself in results.
- If True, includes both stocks and index. If False, returns only stocks.
- Defaults to False.
- Returns:
- list: A list of dictionaries containing stock quote data.
- Each dictionary contains various fields including:
- - symbol: Stock symbol
- - open: Opening price
- - high: High price
- - low: Low price
- - lastPrice: Last traded price
- - change: Change in price
- - pChange: Percentage change
- And other relevant trading information.
- Raises:
- URLError: If unable to connect to NSE servers
- ValueError: If invalid index name is provided
- Example:
- >>> nse = Nse()
- >>> nifty_quotes = nse.get_stock_quote_in_index("NIFTY 50")
- >>> nifty_quotes_with_index = nse.get_stock_quote_in_index("NIFTY 50", include_index=True)
- """
-
- index = index.upper()
- url = STOCKS_IN_INDEX_URL % index
- res = self.session.fetch(url)
- res_dict = res.json()
- res_dict = cast_intfloat_string_values_to_intfloat(res_dict)
- if include_index is False:
- return [record for record in res_dict['data'] if record['priority'] == 0]
- else:
- return res_dict['data']
-
- def _get_top_gainers_losers(self, direction, index):
- """Internal method to fetch top gainers or losers for a given index.
-
- Args:
- direction (str): Either 'gainers' or 'losers'
- index (str): Index name - one of NIFTY, BANKNIFTY, NIFTYNEXT50, SecGtr20, SecLwr20, FNO, ALL
-
- Returns:
- list: List of dictionaries containing top gainers/losers data for the specified index
-
- Raises:
- ValueError: If invalid index name is provided
- """
- index = index or 'NIFTY' # Default to NIFTY if None
- index = index.upper()
- index = {
- "NIFTY": "NIFTY",
- "NIFTY 50": "NIFTY",
- "NIFTY BANK": "BANKNIFTY",
- "BANKNIFTY": "BANKNIFTY",
- "NIFTYNEXT50": "NIFTYNEXT50",
- "NIFTY NEXT 50": "NIFTYNEXT50",
- "SECGTR20": "SecGtr20",
- "SECLWR20": "SecLwr20",
- "FNO": "FOSec",
- "ALL": "allSec"
- }.get(index)
- if index is None:
- raise ValueError("Index must be one of NIFTY 50, NIFTY BANK, NIFTY NEXT 50, SecGtr20, SecLwr20, FNO, ALL")
- url = TOP_GAINERS_URL if direction == 'gainers' else TOP_LOSERS_URL
- res = self.session.fetch(url)
- return cast_intfloat_string_values_to_intfloat(res.json())[index]['data']
-
- #############################
- ### DERIVATIVE APIS ###
- #############################
-
- def get_future_quote(self, code, expiry_date=None):
- """Get future quote for given stock code.
-
- This function fetches futures trading data for a given stock code from NSE's derivatives segment.
- If expiry date is provided, returns data for that specific expiry, else returns data for all
- available expiry dates.
-
- Args:
- code (str): Stock code for which futures data needs to be fetched
- expiry_date (str, optional): Expiry date in format DD-MMM-YYYY (e.g. "27-Mar-2025").
- Defaults to None.
-
- Returns:
- Union[dict, list]: If expiry_date provided returns dict with futures data for that expiry,
- else returns list of dicts with data for all expiries.
-
- Example:
- >>> nse = Nse()
- >>> nse.get_future_quote('RELIANCE')
- [{'expiryDate': '27-Mar-2025',
- 'lastPrice': 1246,
- 'premium': 4.45,
- 'openPrice': 1245.25,
- 'highPrice': 1260.85,
- 'lowPrice': 1236.2,
- 'openInterest': 257812,
- 'changeInOpenInterest': 7144,
- ...},
- {...}]
- """
-
- url = QUOTE_DRIVATIVE_URL % code.upper()
- res = self.session.fetch(url)
- res_dict = res.json()
- # list containing all options and futures data
- data = res_dict['stocks']
- # filter out only future data
- future_data = [s for s in data if s['metadata']['instrumentType'] == "Stock Futures"]
- # future data is very convoluted, so flatten-out the desired data
- # !! there is bug in spelling of the key 'dailyvolatility', it is not camel cased
- # fixing that in my code for uniformity
- filtered_data = [
- {
- 'expiryDate': record['metadata']['expiryDate'],
- 'lastPrice': record['metadata']['lastPrice'],
- 'premium': record['metadata']['lastPrice'] - record['underlyingValue'],
- 'openPrice': record['metadata']['openPrice'],
- 'highPrice': record['metadata']['highPrice'],
- 'lowPrice': record['metadata']['lowPrice'],
- 'closePrice': record['metadata']['closePrice'],
- 'prevClose': record['metadata']['prevClose'],
- 'change': record['metadata']['change'],
- 'pChange': record['metadata']['pChange'],
- 'numberOfContractsTraded': record['metadata']['numberOfContractsTraded'],
- 'totalTurnover': record['metadata']['totalTurnover'],
- 'underlyingValue': record['underlyingValue'],
- 'tradedVolume': record['marketDeptOrderBook']['tradeInfo']['tradedVolume'],
- 'openInterest': record['marketDeptOrderBook']['tradeInfo']['openInterest'],
- 'changeInOpenInterest': record['marketDeptOrderBook']['tradeInfo']['changeinOpenInterest'],
- 'pchangeinOpenInterest': record['marketDeptOrderBook']['tradeInfo']['pchangeinOpenInterest'],
- 'marketLot': record['marketDeptOrderBook']['tradeInfo']['marketLot'],
- 'dailyVolatility': record['marketDeptOrderBook']['otherInfo']['dailyvolatility'],
- 'annualisedVolatility': record['marketDeptOrderBook']['otherInfo']['annualisedVolatility']
- }
- for record in future_data
- ]
- # if expiry_date is provided, filter out data for that expiry date
- if expiry_date:
- matching_records = [record for record in filtered_data if record['expiryDate'] == expiry_date]
- if matching_records:
- return matching_records[0] # Return the first matching record
- else:
- # Return an empty dictionary if no records found for the given expiry date
- return {}
- return filtered_data
-
- def __str__(self):
- """Returns a string representation of the NSE driver class.
- Returns:
- str: A descriptive string identifying this as the NSE driver class.
- """
-
- return 'Driver Class for National Stock Exchange (NSE)'
-
-
-if __name__ == "__main__":
- n = Nse()
- # data = n.download_bhavcopy("14th Dec")
- n.get_quote('reliance')
diff --git a/singular_ticker_causal/data_sources/nseconnect/ua.py b/singular_ticker_causal/data_sources/nseconnect/ua.py
deleted file mode 100644
index 2095150b263678dc19f3f3c2eebd2c5b49a16eb7..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/data_sources/nseconnect/ua.py
+++ /dev/null
@@ -1,137 +0,0 @@
-import requests
-import random
-from datetime import datetime as dt
-from .urls import NSE_MAIN
-from time import sleep
-
-
-class Session():
- __CACHE__ = {}
-
- def __init__(self, session_refresh_interval=60, cache_timeout=60):
- """Initialize the class instance with session and cache parameters.
- Args:
- session_refresh_interval (int, optional): Time interval in seconds to refresh session. Defaults to 60.
- cache_timeout (int, optional): Cache timeout duration in seconds. Defaults to 20.
- Attributes:
- session_refresh_interval (int): Time interval for session refresh.
- cache_timeout (int): Duration for cache timeout.
- """
-
- self.session_refresh_interval = session_refresh_interval
- self.cache_timeout = cache_timeout # cache timeout in seconds
- self._session = None # Initialize _session attribute to None
- self.create_session()
- self.flush()
-
- def nse_headers(self):
- """Returns a dictionary of headers required for making requests to NSE (National Stock Exchange).
- These headers are designed to mimic a web browser request to prevent request blocking.
- Returns:
- dict: A dictionary containing HTTP headers with the following keys:
- - Accept: Acceptable content types
- - Accept-Language: Preferred language for response
- - user-agent: Browser identification string
- - X-Requested-With: Identifies AJAX requests
- """
-
- return {
- "Accept": "text/html,application/xhtml+xml,application/xml;q=0.9,image/avif,image/webp,image/apng,*/*;q=0.8,application/signed-exchange;v=b3;q=0.7",
- "Accept-Language": "en-US,en;q=0.9",
- "user-agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/122.0.0.0 Safari/537.36",
- "X-Requested-With": "XMLHttpRequest",
- "Referer": "https://www.nseindia.com/",
- "Origin": "https://www.nseindia.com",
- "Connection": "keep-alive",
- "Sec-Fetch-Dest": "empty",
- "Sec-Fetch-Mode": "cors",
- "Sec-Fetch-Site": "same-origin"
- }
-
- def create_session(self):
- """Creates and initializes a new HTTP session for NSE (National Stock Exchange) API requests.
- This method sets up a requests.Session object with appropriate headers for NSE and initializes
- it by making a GET request to the NSE home page. The session is used for subsequent API calls.
- Returns:
- None
- Side Effects:
- - Sets self._session with configured requests.Session object
- - Sets self._session_init_time with current timestamp
- """
-
- # Clean up old session if it exists
- if hasattr(self, '_session') and self._session is not None:
- old_session = self._session
- self._session = None
- # Explicitly delete old session to ensure garbage collection
- del old_session
-
- # Create a completely new session object
- self._session = requests.Session()
- self._session.headers.update(self.nse_headers())
-
- # First visit NSE home page to get cookies
- self._session.get(NSE_MAIN)
- # Small delay to mimic human behavior
- sleep(1)
- # Visit the market page to get additional cookies
- self._session.get(f"{NSE_MAIN}/market-data/live-equity-market")
-
- self._session_init_time = dt.now()
-
- def flush(self):
- """Flushes the cached user agent data.
- This method clears the internal cache dictionary storing user agent information
- by resetting the class's __CACHE__ attribute to an empty dictionary.
- Returns:
- None
- """
-
- self.__class__.__CACHE__ = {}
-
- def fetch(self, url):
- """Fetches data from a given URL with caching and session management.
- This method implements a caching mechanism and session refresh logic to optimize
- network requests. It also includes random delays to prevent rate limiting.
- Args:
- url (str): The URL to fetch data from.
- Returns:
- requests.Response: The response object from the request.
- Note:
- - Uses class-level cache to store responses
- - Implements random delays between 0-300ms before making requests
- - Auto-refreshes session if expired based on session_refresh_interval
- """
-
- # Check cache first
- if url in self.__class__.__CACHE__:
- cache_time, response = self.__class__.__CACHE__[url]
- if (dt.now() - cache_time).seconds < self.cache_timeout:
- # print("serving from cache")
- return response
-
- # Only check session expiry if we need to make a network request
- time_diff = dt.now() - self._session_init_time
- if time_diff.seconds >= self.session_refresh_interval:
- # print("re-initing the session because of expiry")
- self.create_session()
-
- # Add random delay before making request
- sleep_time = random.uniform(0, 0.3) # Random delay between 0-300ms
- # print(f"Adding random delay of {sleep_time:.3f} seconds")
- sleep(sleep_time)
-
- # Make actual request if not in cache or cache expired
- try:
- response = self._session.get(url)
- # Force a 401 response to retry with a fresh session
- if response.status_code == 401:
- self.create_session()
- response = self._session.get(url)
- except requests.RequestException:
- # Try again with a fresh session on any request exception
- self.create_session()
- response = self._session.get(url)
-
- self.__class__.__CACHE__[url] = (dt.now(), response)
- return response
diff --git a/singular_ticker_causal/data_sources/nseconnect/urls.py b/singular_ticker_causal/data_sources/nseconnect/urls.py
deleted file mode 100644
index 9798de352cd31d1c9f47fb6bea52512d66ddb6ae..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/data_sources/nseconnect/urls.py
+++ /dev/null
@@ -1,35 +0,0 @@
-"""
-URL constants for NSE related operations
-"""
-
-# Base URLs
-NSE_HOME = "https://www.nseindia.com"
-NSE_MAIN = "https://www.nseindia.com"
-NSE_LEGACY = "https://www.nseindia.com"
-
-# Quote URLs
-QUOTE_EQUITY_URL = f"{NSE_MAIN}/get-quotes/equity?symbol=%s"
-QUOTE_API_URL = f"{NSE_MAIN}/api/quote-equity?symbol=%s"
-
-# Stock list URLs
-STOCKS_CSV_URL = f"https://archives.nseindia.com/content/equities/EQUITY_L.csv"
-
-# Market movers URLs
-TOP_GAINERS_URL = f"{NSE_MAIN}/api/live-analysis-variations?index=gainers"
-TOP_LOSERS_URL = f"{NSE_MAIN}/api/live-analysis-variations?index=loosers"
-TOP_FNO_GAINER_URL = f"{NSE_MAIN}/api/market-data-pre-open?key=FO"
-TOP_FNO_LOSER_URL = f"{NSE_MAIN}/api/market-data-pre-open?key=FO"
-FIFTYTWO_WEEK_HIGH_URL = f"{NSE_MAIN}/api/live-analysis-52Week?index=high"
-FIFTYTWO_WEEK_LOW_URL = f"{NSE_MAIN}/api/live-analysis-52Week?index=low"
-
-# Index URLs
-ALL_INDICES_URL = f"{NSE_MAIN}/api/allIndices"
-STOCKS_IN_INDEX_URL = f"{NSE_MAIN}/api/equity-stockIndices?index=%s"
-
-
-# Historical data URLs
-BHAVCOPY_BASE_URL = f"{NSE_MAIN}/archives/equities-bhavcopy/%s"
-BHAVCOPY_BASE_FILENAME = "cm%s%s%sbhav.csv"
-
-# Drivative URLs
-QUOTE_DRIVATIVE_URL = f"{NSE_MAIN}/api/quote-derivative?symbol=%s"
diff --git a/singular_ticker_causal/data_sources/nseconnect/utils.py b/singular_ticker_causal/data_sources/nseconnect/utils.py
deleted file mode 100644
index e670b7863cd2bc93c108be98d6a1ce39d7e620ab..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/data_sources/nseconnect/utils.py
+++ /dev/null
@@ -1,373 +0,0 @@
-"""
- The MIT License (MIT)
-
- Copyright (c) 2014 Noufal Nazar
-
- Permission is hereby granted, free of charge, to any person obtaining a copy
- of this software and associated documentation files (the "Software"), to deal
- in the Software without restriction, including without limitation the rights
- to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
- copies of the Software, and to permit persons to whom the Software is
- furnished to do so, subject to the following conditions:
-
- The above copyright notice and this permission notice shall be included in all
- copies or substantial portions of the Software.
-
- THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
- IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
- FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
- AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
- LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
- OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
- SOFTWARE.
-
-"""
-import six
-import re
-import operator
-
-def byte_adaptor(fbuffer):
- """ provides py3 compatibility by converting byte based
- file stream to string based file stream
-
- Arguments:
- fbuffer: file like objects containing bytes
-
- Returns:
- string buffer
- """
- if six.PY3:
- strings = fbuffer.read().decode('latin-1')
- fbuffer = six.StringIO(strings)
- return fbuffer
- else:
- return fbuffer
-
-
-def js_adaptor(buffer):
- """ convert javascript objects like true, none, NaN etc. to
- quoted word.
-
- Arguments:
- buffer: string to be converted
-
- Returns:
- string after conversion
- """
- buffer = re.sub('true', 'True', buffer)
- buffer = re.sub('false', 'False', buffer)
- buffer = re.sub('none', 'None', buffer)
- buffer = re.sub('NaN', '"NaN"', buffer)
- return buffer
-
-def cast_intfloat_string_values_to_intfloat(data, round_digits=2):
- """Recursively converts string representations of numbers to integers or floats in nested data structures.
- This function traverses through dictionaries and lists, converting string values that represent
- numbers into their corresponding numeric types (int or float). For float values, it rounds to
- the specified number of decimal places.
- Args:
- data (Union[dict, list]): The input data structure containing values to be converted.
- Can be either a dictionary or a list, potentially nested.
- round_digits (int, optional): Number of decimal places to round float values to.
- Defaults to 2.
- Returns:
- Union[dict, list]: A new data structure of the same type as input, with string
- representations of numbers converted to their numeric types.
- Example:
- >>> data = {'a': '1', 'b': '2.5', 'c': 'text', 'd': {'e': '3.14'}}
- >>> cast_intfloat_string_values_to_intfloat(data)
- {'a': 1, 'b': 2.5, 'c': 'text', 'd': {'e': 3.14}}
- """
-
- if isinstance(data, dict):
- data = data.copy()
- for key, value in data.items():
- if isinstance(value, str):
- try:
- data[key] = int(value)
- except ValueError:
- try:
- data[key] = round(float(value), round_digits)
- except ValueError:
- pass
- elif isinstance(value, (dict, list)):
- data[key] = cast_intfloat_string_values_to_intfloat(value, round_digits)
- elif isinstance(value, float):
- data[key] = round(value, round_digits)
- elif isinstance(data, list):
- data = data[:]
- for i, value in enumerate(data):
- if isinstance(value, str):
- try:
- data[i] = int(value)
- except ValueError:
- try:
- data[i] = round(float(value), round_digits)
- except ValueError:
- pass
- elif isinstance(value, (dict, list)):
- data[i] = cast_intfloat_string_values_to_intfloat(value, round_digits)
- elif isinstance(value, float):
- data[i] = round(value, round_digits)
- return data
-
-def camel_to_title(camel_str):
- """Converts a camel case string to title case.
- This function takes a camel case string and converts it to title case by adding
- spaces before capital letters and capitalizing the first letter of each word.
- Args:
- camel_str (str): The camel case string to be converted.
- Returns:
- str: The converted string in title case format.
- Examples:
- >>> camel_to_title("camelCaseString")
- 'Camel Case String'
- >>> camel_to_title("thisIsATest")
- 'This Is A Test'
- """
-
- return re.sub(r'(?=': operator.ge,
- '<=': operator.le,
- '>': operator.gt,
- '<': operator.lt
- }
-
- for op_str, op_func in operators.items():
- if op_str in query_str:
- path, value = query_str.split(op_str)
- path = path.strip()
- value = value.strip()
-
- # Try to convert value to number if possible
- try:
- value = int(value)
- except ValueError:
- try:
- value = float(value)
- except ValueError:
- # Keep as string if not numeric
- pass
-
- return path, op_func, value
-
- return None, None, None
-
-def dict_to_table(data, title="Data Table", filter=None, ignore=None, sort=None, direction="desc", query=None):
- """Converts dictionary or list of dictionaries to a formatted table using Rich library.
- This function takes either a dictionary or a list of dictionaries and displays it as a
- formatted table in the console. It supports filtering specific keys, ignoring keys, and
- applies special formatting for negative numbers.
- Args:
- data (Union[dict, List[dict]]): The data to be displayed. Can be either a dictionary
- or a list of dictionaries.
- title (str, optional): The title to display above the table. Defaults to "Data Table".
- filter (List[str], optional): List of keys to include in the output. If provided,
- only these keys will be displayed. Keys are matched case-insensitively.
- Defaults to None.
- ignore (List[str], optional): List of keys to exclude from the output. Keys are
- matched case-insensitively. Defaults to None.
- sort (str, optional): Key to sort by. Case-insensitive. Will sort numerically
- for numeric values and alphabetically for string values. Defaults to None.
- direction (str, optional): Sort direction - "asc" for ascending or "desc" for
- descending. Defaults to "desc".
- query (str, optional): Filter rows using dot notation path and comparison.
- Supports operators: ==, !=, >, <, >=, <=
- Example: "market.price>100" or "status.active==True"
- Keys are matched case-insensitively. Defaults to None.
- """
- from rich.console import Console
- from rich.table import Table
-
- console = Console()
- table = Table(title=title)
-
- if not data:
- console.print("[red]No data to display![/red]")
- return
-
- # Parse query if provided
- query_path = None
- query_op = None
- query_value = None
- if query:
- query_path, query_op, query_value = _parse_query(query)
- if not all([query_path, query_op, query_value]):
- console.print("[red]Invalid query format![/red]")
- return
-
- # Validate direction
- if direction not in ["asc", "desc"]:
- console.print("[red]Direction must be 'asc' or 'desc'![/red]")
- return
-
- # Normalize filter, ignore and sort keys
- if filter:
- if not isinstance(filter, list):
- console.print("[red]Filter should be a list of keys![/red]")
- return
- filter = [str(key).lower() for key in filter]
-
- if ignore:
- if not isinstance(ignore, list):
- console.print("[red]Ignore should be a list of keys![/red]")
- return
- ignore = [str(key).lower() for key in ignore]
- else:
- ignore = []
-
- if sort:
- sort = str(sort).lower()
-
- # Check if data is a list of dicts
- if isinstance(data, list) and all(isinstance(i, dict) for i in data):
- # Get all unique keys and create key mapping
- keys = set()
- for item in data:
- keys.update(item.keys())
- key_map = {k.lower(): k for k in keys}
-
- # Validate sort key if provided
- if sort and sort not in key_map:
- console.print(f"[red]Sort key '{sort}' not found in data![/red]")
- return
-
- # Create ordered keys list
- if filter:
- ordered_keys = [key_map[f] for f in filter if f in key_map and f not in ignore]
- else:
- ordered_keys = [key_map[k.lower()] for k in keys if k.lower() not in ignore]
-
- if not ordered_keys:
- console.print("[red]No matching keys found![/red]")
- return
-
- # Apply query filter before sorting
- if query:
- filtered_data = []
- for item in data:
- item_value = _resolve_path(item, query_path)
- if item_value is not None:
- try:
- if query_op(item_value, query_value):
- filtered_data.append(item)
- except TypeError:
- # Handle type mismatch gracefully
- continue
- data = filtered_data
-
- if not data:
- console.print("[red]No data matches the query![/red]")
- return
-
- # Sort data if sort key is provided
- if sort and sort in key_map:
- original_key = key_map[sort]
- # Try numeric sort first
- try:
- sorted_data = sorted(
- data,
- key=lambda x: float(x.get(original_key, 0)),
- reverse=(direction == "desc")
- )
- except (ValueError, TypeError):
- # Fall back to string sort
- sorted_data = sorted(
- data,
- key=lambda x: str(x.get(original_key, "")),
- reverse=(direction == "desc")
- )
- else:
- sorted_data = data
-
- # Add columns and display table
- for key in ordered_keys:
- table.add_column(camel_to_title(key), style="bright_white")
-
- for item in sorted_data:
- row = []
- for key in ordered_keys:
- value = item.get(key, "")
- if isinstance(value, (int, float)) and value < 0:
- row.append(f"[red]{value}[/red]")
- else:
- row.append(f"[bright_white]{value}[/bright_white]")
- table.add_row(*row)
-
- elif isinstance(data, dict):
- # Single dict can't be queried for rows
- if query:
- console.print("[red]Query is only supported for list of dictionaries![/red]")
- return
-
- # Filter and ignore the dictionary data
- filtered_data = {}
- key_map = {k.lower(): k for k in data.keys()}
-
- if filter:
- # Add keys in filter order if they exist and not in ignore
- for f in filter:
- if f in key_map and f not in ignore:
- original_key = key_map[f]
- value = data[original_key]
- if not isinstance(value, (dict, list, tuple, set)):
- filtered_data[original_key] = value
- else:
- # If no filter, exclude ignored and nested items
- filtered_data = {k: v for k, v in data.items()
- if not isinstance(v, (dict, list, tuple, set))
- and k.lower() not in ignore}
-
- if not filtered_data:
- console.print("[red]No matching key-value pairs to display![/red]")
- return
-
- # Add columns
- table.add_column("Key", style="cyan", no_wrap=True)
- table.add_column("Value", style="bright_white")
-
- # Add rows
- for key, value in filtered_data.items():
- if isinstance(value, (int, float)) and value < 0:
- value_str = f"[red]{value}[/red]"
- else:
- value_str = f"[bright_white]{value}[/bright_white]"
- table.add_row(camel_to_title(key), value_str)
-
- else:
- console.print("[red]Unsupported data format![/red]")
- return
-
- console.print(table)
-
-
-
diff --git a/singular_ticker_causal/data_sources/sebi_reg30_client.py b/singular_ticker_causal/data_sources/sebi_reg30_client.py
deleted file mode 100644
index 810fb7515c9aa564232b5f6a1db0e91878e79f69..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/data_sources/sebi_reg30_client.py
+++ /dev/null
@@ -1,126 +0,0 @@
-import logging
-from typing import List, Dict, Any
-from datetime import datetime
-import pandas as pd
-from nsepython import nsefetch
-
-
-logger = logging.getLogger(__name__)
-
-
-class SEBIREG30Client:
- """
- Fetch SEBI Regulation 30 corporate announcements from NSE.
- """
-
- BASE_URL = (
- "https://www.nseindia.com/api/corporate-announcements"
- )
-
- def __init__(self):
- pass
-
- @staticmethod
- def _format_date(date_str: str) -> str:
- """
- Convert YYYY-MM-DD -> DD-MM-YYYY
- NSE API expects DD-MM-YYYY
- """
- return datetime.strptime(date_str, "%Y-%m-%d").strftime("%d-%m-%Y")
-
- def fetch(
- self,
- ticker: str,
- start: str,
- end: str,
- ) -> List[Dict[str, Any]]:
- """
- Fetch corporate announcements for a ticker
- within a date range.
-
- Args:
- ticker: NSE symbol (e.g. TCS)
- start: YYYY-MM-DD
- end: YYYY-MM-DD
-
- Returns:
- List of announcement dicts
- """
-
- logger.info(
- f"Fetching SEBI Reg 30 announcements for "
- f"{ticker} from {start} to {end}"
- )
-
- try:
- start_fmt = self._format_date(start)
- end_fmt = self._format_date(end)
-
- # NSE announcement endpoint
- url = (
- f"{self.BASE_URL}"
- f"?index=equities"
- f"&symbol={ticker.upper()}"
- f"&from_date={start_fmt}"
- f"&to_date={end_fmt}"
- )
-
- logger.info(f"NSE URL: {url}")
-
- data = nsefetch(url)
-
- if not data:
- logger.warning("No announcement data returned")
- return []
-
- # Normalize into dataframe
- announcements = pd.json_normalize(data)
-
- return announcements
-
- except Exception as e:
- logger.exception(
- f"Error fetching NSE announcements: {e}"
- )
- return []
-
-
-if __name__ == "__main__":
- import json
- from datetime import timedelta
-
- logging.basicConfig(
- level=logging.INFO,
- format=(
- "%(asctime)s - %(name)s - "
- "%(levelname)s - %(message)s"
- )
- )
-
- client = SEBIREG30Client()
-
- ticker = "TCS"
-
- end_date = datetime.now().strftime("%Y-%m-%d")
- start_date = (
- datetime.now() - timedelta(days=30)
- ).strftime("%Y-%m-%d")
-
- print(f"\n--- Testing SEBI Reg 30 Fetch: {ticker} ---")
-
- announcements = client.fetch(
- ticker,
- start_date,
- end_date
- )
- announcements.to_csv("sebi.csv")
-
- if announcements:
- print("\nLatest Announcement:\n")
- print(json.dumps(
- announcements[0],
- indent=4,
- default=str
- ))
- else:
- print("No announcements found.")
\ No newline at end of file
diff --git a/singular_ticker_causal/data_sources/test.py b/singular_ticker_causal/data_sources/test.py
deleted file mode 100644
index dc402cc66e361bb9a984e5c0100cda09757534f2..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/data_sources/test.py
+++ /dev/null
@@ -1,7 +0,0 @@
-import yfinance as yf
-
-ticker = yf.Ticker("INFY.NS")
-print(dir(ticker))
-
-
-# Output - ['__class__', '__delattr__', '__dict__', '__dir__', '__doc__', '__eq__', '__format__', '__ge__', '__getattribute__', '__getstate__', '__gt__', '__hash__', '__init__', '__init_subclass__', '__le__', '__lt__', '__module__', '__ne__', '__new__', '__reduce__', '__reduce_ex__', '__repr__', '__setattr__', '__sizeof__', '__str__', '__subclasshook__', '__weakref__', '_analysis', '_data', '_download_options', '_earnings', '_earnings_dates', '_expirations', '_fast_info', '_fetch_ticker_tz', '_financials', '_fundamentals', '_funds_data', '_get_earnings_dates_using_scrape', '_get_earnings_dates_using_screener', '_get_ticker_tz', '_holders', '_isin', '_lazy_load_price_history', '_message_handler', '_news', '_options2df', '_price_history', '_quote', '_shares', '_tz', '_underlying', 'actions', 'analyst_price_targets', 'balance_sheet', 'balancesheet', 'calendar', 'capital_gains', 'cash_flow', 'cashflow', 'dividends', 'earnings', 'earnings_dates', 'earnings_estimate', 'earnings_history', 'eps_revisions', 'eps_trend', 'fast_info', 'financials', 'funds_data', 'get_actions', 'get_analyst_price_targets', 'get_balance_sheet', 'get_balancesheet', 'get_calendar', 'get_capital_gains', 'get_cash_flow', 'get_cashflow', 'get_dividends', 'get_earnings', 'get_earnings_dates', 'get_earnings_estimate', 'get_earnings_history', 'get_eps_revisions', 'get_eps_trend', 'get_fast_info', 'get_financials', 'get_funds_data', 'get_growth_estimates', 'get_history_metadata', 'get_income_stmt', 'get_incomestmt', 'get_info', 'get_insider_purchases', 'get_insider_roster_holders', 'get_insider_transactions', 'get_institutional_holders', 'get_isin', 'get_major_holders', 'get_mutualfund_holders', 'get_news', 'get_recommendations', 'get_recommendations_summary', 'get_revenue_estimate', 'get_sec_filings', 'get_shares', 'get_shares_full', 'get_splits', 'get_sustainability', 'get_upgrades_downgrades', 'get_valuation_measures', 'growth_estimates', 'history', 'history_metadata', 'income_stmt', 'incomestmt', 'info', 'insider_purchases', 'insider_roster_holders', 'insider_transactions', 'institutional_holders', 'isin', 'live', 'major_holders', 'mutualfund_holders', 'news', 'option_chain', 'options', 'quarterly_balance_sheet', 'quarterly_balancesheet', 'quarterly_cash_flow', 'quarterly_cashflow', 'quarterly_earnings', 'quarterly_financials', 'quarterly_income_stmt', 'quarterly_incomestmt', 'recommendations', 'recommendations_summary', 'revenue_estimate', 'sec_filings', 'session', 'shares', 'splits', 'sustainability', 'ticker', 'ttm_cash_flow', 'ttm_cashflow', 'ttm_financials', 'ttm_income_stmt', 'ttm_incomestmt', 'upgrades_downgrades', 'valuation', 'ws']
\ No newline at end of file
diff --git a/singular_ticker_causal/services/__init__.py b/singular_ticker_causal/services/__init__.py
deleted file mode 100644
index 1491875195b619cb835bb692851dfe69b70dcbdd..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/services/__init__.py
+++ /dev/null
@@ -1,3 +0,0 @@
-from .pelt_detection import PELTDetector as BOCDDetector
-from .ssa_denoiser import MSSAQuantEngine
-from .tensor_builder import TensorBuilder, EmbeddingService
\ No newline at end of file
diff --git a/singular_ticker_causal/services/camef_simulator.py b/singular_ticker_causal/services/camef_simulator.py
deleted file mode 100644
index 7263fb29f883e9e6f0859a3be169b48b7c6ed7cd..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/services/camef_simulator.py
+++ /dev/null
@@ -1,420 +0,0 @@
-"""
-camef_simulator.py
-
-CAMEF (Causal-Augmented Multi-Modality Event-Driven Financial Forecasting)
-Stress-test wrapper around GPT4MTS for counterfactual scenario evaluation
-(e.g. "RBI cuts interest rates by 50bps").
-
-Modes (priority order):
- 1. checkpoint – GPT4MTS loaded from a trained .pth checkpoint.
- All backbone weights come from the checkpoint.
- 2. zero-shot – GPT4MTS initialised from pre-trained HF hub weights
- (RoBERTa-base + GPT-2 + MOMENT-1-large) with no fine-tuning.
- Output is uncalibrated but architecturally correct.
-
-Key tensor conventions (matching GPT4MTS.predict_single_case):
- batch_seq : (batch=1, seq_len, d) ← time-series input
- output : (batch=1, d, pred_len) ← raw model output
- returned : (1, d, pred_len) ← what simulate_shock returns
-"""
-
-import logging
-from typing import Optional
-import torch
-from singular_ticker_causal.algorithms.CAMEF.CAMEF import GPT4MTS
-
-
-logger = logging.getLogger(__name__)
-
-
-# ── Public simulator ─────────────────────────────────────────────────────────
-
-class CAMEFSimulator:
- """
- Thin wrapper around GPT4MTS for stress-test inference.
-
- Parameters
- ----------
- model_path : str, optional
- Path to a .pth checkpoint saved by GPT4MTS.save_model_combined().
- If None or loading fails, falls back to zero-shot HF weights.
- device : str
- 'cuda' or 'cpu'.
- seq_len : int
- Number of historical time-steps fed to the model. Must match the
- value used during training (or be consistent for zero-shot).
- pred_len : int
- Forecast horizon. Drives output_project output size.
- d : int
- Number of variates per time-step (1 for univariate nodes).
- bert : str
- HF model id for the RoBERTa text encoder.
- gpt : str
- HF model id for the GPT-2 decoder.
- moment : str
- HF model id for the MOMENT time-series encoder (zero-shot mode only).
- Ignored when loading from a checkpoint.
- window : int
- Sliding-window size for the RoBERTa tokeniser.
- stride : int
- Stride for the RoBERTa sliding window.
- """
-
- def __init__(
- self,
- model_path: Optional[str] = None,
- device: str = "cpu",
- seq_len: int = 10,
- pred_len: int = 5,
- d: int = 1,
- bert: str = "roberta-base",
- gpt: str = "gpt2",
- moment: str = "AutonLab/MOMENT-1-large",
- window: int = 512,
- stride: int = 256,
- ):
- self.device = torch.device(device if torch.cuda.is_available() else "cpu")
- self.seq_len = seq_len
- self.pred_len = pred_len
- self.d = d
- self._mode = None
- self._model: Optional[GPT4MTS] = None
-
- # ── Attempt 1: load trained checkpoint ───────────────────────────────
- if model_path:
- try:
- logger.info(f"CAMEFSimulator: loading checkpoint from {model_path} ...")
- m = GPT4MTS(
- bert=bert,
- moment=moment,
- gpt=gpt,
- seq_len=seq_len,
- pred_len=pred_len,
- d=d,
- window=window,
- stride=stride,
- batch_size=1,
- )
- m.load_model_combined(save_path=model_path)
- m.to(self.device)
- m.eval()
- self._model = m
- self._mode = "checkpoint"
- logger.info("CAMEFSimulator: checkpoint mode active.")
- except Exception as exc:
- logger.warning(
- f"CAMEFSimulator: checkpoint load failed ({exc}). "
- "Falling back to zero-shot mode."
- )
-
- # ── Attempt 2: zero-shot with pre-trained HF weights ─────────────────
- if self._model is None:
- try:
- logger.info(
- "CAMEFSimulator: initialising zero-shot GPT4MTS "
- f"(bert={bert}, gpt={gpt}, moment={moment}) ..."
- )
- # NOTE: We do NOT set local_files_only here so that the MOMENT
- # model can be downloaded from the HF hub on first run.
- m = _build_zero_shot_gpt4mts(
- bert=bert,
- gpt=gpt,
- moment=moment,
- seq_len=seq_len,
- pred_len=pred_len,
- d=d,
- window=window,
- stride=stride,
- device=self.device,
- )
- m.eval()
- self._model = m
- self._mode = "zero-shot"
- logger.info(
- "CAMEFSimulator: zero-shot mode active. "
- "Output is uncalibrated — supply a model_path to enable checkpoint mode."
- )
- except Exception as exc:
- logger.error(
- "CAMEFSimulator: zero-shot init failed (%s). "
- "Falling back to deterministic stub mode.",
- exc,
- )
- self._model = _StubGPT4MTS(pred_len=self.pred_len, d=self.d)
- self._mode = "stub"
- logger.warning(
- "CAMEFSimulator: stub mode active. "
- "Outputs are deterministic placeholders for pipeline continuity."
- )
-
- if self._model is None:
- raise RuntimeError("CAMEFSimulator failed to initialize: No model loaded.")
-
- # ── Public API ────────────────────────────────────────────────────────────
-
- def simulate_shock(
- self,
- textual_event: str,
- historical_series: torch.Tensor,
- ) -> torch.Tensor:
- """
- Simulate the causal impact of a textual shock on a historical series.
-
- Parameters
- ----------
- textual_event : str
- Natural-language description of the shock, e.g.
- "RBI cuts interest rates by 50 bps".
- historical_series : torch.Tensor, shape (T,) or (T, d)
- Historical values for one (or more) node(s). Will be
- truncated / padded to seq_len automatically.
-
- Returns
- -------
- torch.Tensor, shape (1, d, pred_len)
- Counterfactual forecast under the shock scenario.
- """
-
- # ── Shape normalisation ───────────────────────────────────────────────
- series = historical_series.float()
- if series.dim() == 1:
- series = series.unsqueeze(-1) # (T,) → (T, d)
-
- # Pad or truncate to seq_len
- T = series.shape[0]
- if T < self.seq_len:
- pad = torch.zeros(self.seq_len - T, self.d)
- series = torch.cat([pad, series], dim=0)
- else:
- series = series[-self.seq_len:] # (seq_len, d)
-
- # GPT4MTS.predict_single_case expects (batch, seq_len, d)
- batch_seq = series.unsqueeze(0).to(self.device) # (1, seq_len, d)
-
- with torch.no_grad():
- # Returns (batch=1, d, pred_len), plus 3 intermediate tensors
- output, _, _, _ = self._model.predict_single_case(
- [textual_event], batch_seq
- )
- return output # (1, d, pred_len)
-
- @property
- def mode(self) -> str:
- """Active mode: 'checkpoint' or 'zero-shot'."""
- return self._mode
-
-
-# ── Helper: build GPT4MTS without local_files_only ───────────────────────────
-
-def _build_zero_shot_gpt4mts(
- bert: str,
- gpt: str,
- moment: str,
- seq_len: int,
- pred_len: int,
- d: int,
- window: int,
- stride: int,
- device: torch.device,
-) -> GPT4MTS:
- """
- Construct a GPT4MTS instance using HF hub weights (no local checkpoint).
- """
- model = GPT4MTS(
- bert=bert,
- moment=moment,
- gpt=gpt,
- seq_len=seq_len,
- pred_len=pred_len,
- d=d,
- window=window,
- stride=stride,
- batch_size=1,
- )
- model.to(device)
- return model
-
-
-# ── Helper: deterministic fallback used when model init fails ───────────────
-
-class _StubGPT4MTS:
- """Minimal drop-in for predict_single_case used to keep tests runnable."""
-
- def __init__(self, pred_len: int, d: int):
- self.pred_len = pred_len
- self.d = d
-
- def predict_single_case(self, _texts, batch_seq: torch.Tensor):
- # batch_seq shape: (1, seq_len, d)
- last = batch_seq[:, -1, :] # (1, d)
- output = last.unsqueeze(-1).repeat(1, 1, self.pred_len) # (1, d, pred_len)
- return output, None, None, None
-
-
-# ── Main: self-contained test harness ────────────────────────────────────────
-
-if __name__ == "__main__":
- import sys
- import logging
-
- logging.basicConfig(
- level=logging.INFO,
- format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
- )
-
- # ── Shared hyper-parameters ───────────────────────────────────────────────
- SEQ_LEN = 10 # keep small – we are not training
- PRED_LEN = 5
- D = 1 # variates per time-step
-
- # ── Helper ────────────────────────────────────────────────────────────────
- def _sep(title: str) -> None:
- print(f"\n{'=' * 62}\n {title}\n{'=' * 62}")
-
- def _ok(msg: str) -> None:
- print(f" ✓ {msg}")
-
- def _fail(msg: str) -> None:
- print(f" ✗ {msg}")
- sys.exit(1)
-
- # ══════════════════════════════════════════════════════════════════════════
- # TEST 1 – Initialisation (zero-shot, no checkpoint path)
- # ══════════════════════════════════════════════════════════════════════════
- _sep("TEST 1 – Zero-shot initialisation (no checkpoint)")
- try:
- simulator = CAMEFSimulator(
- model_path=None,
- device="cpu",
- seq_len=SEQ_LEN,
- pred_len=PRED_LEN,
- d=D,
- )
- except RuntimeError as exc:
- _fail(f"CAMEFSimulator init raised RuntimeError: {exc}")
-
- if simulator.mode not in {"checkpoint", "zero-shot"}:
- _fail(f"Unexpected mode: {simulator.mode!r}")
- _ok(f"Simulator initialised in {simulator.mode!r} mode")
-
- # ══════════════════════════════════════════════════════════════════════════
- # TEST 2 – Univariate simulate_shock: output shape & dtype
- # ══════════════════════════════════════════════════════════════════════════
- _sep("TEST 2 – simulate_shock(): shape & dtype (univariate, d=1)")
-
- torch.manual_seed(0)
- hist_1d = torch.randn(SEQ_LEN) # shape (T,) – 1-D shorthand
- event = "RBI cuts interest rates by 50 bps"
-
- out = simulator.simulate_shock(event, hist_1d)
-
- expected_shape = (1, D, PRED_LEN)
- if out.shape != torch.Size(expected_shape):
- _fail(f"Expected shape {expected_shape}, got {tuple(out.shape)}")
- _ok(f"Output shape : {tuple(out.shape)} (correct)")
-
- if out.dtype != torch.float32:
- _fail(f"Expected float32, got {out.dtype}")
- _ok(f"Output dtype : {out.dtype} (correct)")
-
- if torch.isnan(out).any() or torch.isinf(out).any():
- _fail("Output contains NaN / Inf values")
- _ok("Output is finite (no NaN / Inf)")
-
- print(f" Forecast values: {out.squeeze().detach().tolist()}")
-
- # ══════════════════════════════════════════════════════════════════════════
- # TEST 3 – Multivariate simulate_shock (d=2)
- # ══════════════════════════════════════════════════════════════════════════
- _sep("TEST 3 – simulate_shock(): multivariate series (d=2)")
- D2 = 2
- sim2 = CAMEFSimulator(
- model_path=None,
- device="cpu",
- seq_len=SEQ_LEN,
- pred_len=PRED_LEN,
- d=D2,
- )
- hist_2d = torch.randn(SEQ_LEN, D2) # shape (T, d)
- out2 = sim2.simulate_shock(event, hist_2d)
-
- expected_shape2 = (1, D2, PRED_LEN)
- if out2.shape != torch.Size(expected_shape2):
- _fail(f"Expected shape {expected_shape2}, got {tuple(out2.shape)}")
- _ok(f"Output shape : {tuple(out2.shape)} (correct)")
-
- # ══════════════════════════════════════════════════════════════════════════
- # TEST 4 – Edge case: series shorter than seq_len (pad path)
- # ══════════════════════════════════════════════════════════════════════════
- _sep(f"TEST 4 – Edge case: T < seq_len (pad with zeros)")
- T_short = SEQ_LEN - 3 # 3 steps fewer than seq_len
- hist_short = torch.randn(T_short)
- out_short = simulator.simulate_shock(event, hist_short)
-
- if out_short.shape != torch.Size(expected_shape):
- _fail(f"Expected shape {expected_shape}, got {tuple(out_short.shape)}")
- _ok(f"Padding path: output shape {tuple(out_short.shape)} (correct)")
-
- # ══════════════════════════════════════════════════════════════════════════
- # TEST 5 – Edge case: series longer than seq_len (truncation path)
- # ══════════════════════════════════════════════════════════════════════════
- _sep(f"TEST 5 – Edge case: T > seq_len (truncation path)")
- T_long = SEQ_LEN + 20
- hist_long = torch.randn(T_long)
- out_long = simulator.simulate_shock(event, hist_long)
-
- if out_long.shape != torch.Size(expected_shape):
- _fail(f"Expected shape {expected_shape}, got {tuple(out_long.shape)}")
- _ok(f"Truncation path: output shape {tuple(out_long.shape)} (correct)")
-
- # ══════════════════════════════════════════════════════════════════════════
- # TEST 6 – Repeatability: same input → same output (eval / no-grad)
- # ══════════════════════════════════════════════════════════════════════════
- _sep("TEST 6 – Repeatability (eval mode, identical inputs)")
- torch.manual_seed(7)
- hist_rep = torch.randn(SEQ_LEN)
- out_a = simulator.simulate_shock(event, hist_rep)
- out_b = simulator.simulate_shock(event, hist_rep)
-
- if not torch.allclose(out_a, out_b, atol=1e-5):
- _fail("Outputs differ across identical calls (non-determinism detected)")
- _ok("Identical inputs produce identical outputs")
-
- # ══════════════════════════════════════════════════════════════════════════
- # TEST 7 – Scenario comparison: dovish vs hawkish shock
- # Mimics the stress-test workflow used in the causal pipeline:
- # run the same historical window under two contrasting events and confirm
- # that the *direction* of the model's forecast differs.
- # ══════════════════════════════════════════════════════════════════════════
- _sep("TEST 7 – Scenario comparison (dovish vs hawkish shock)")
-
- torch.manual_seed(42)
- hist_base = torch.randn(SEQ_LEN)
-
- event_dovish = "RBI cuts interest rates by 50 bps, signalling accommodative stance"
- event_hawkish = "RBI hikes interest rates by 75 bps to combat elevated inflation"
-
- out_dovish = simulator.simulate_shock(event_dovish, hist_base)
- out_hawkish = simulator.simulate_shock(event_hawkish, hist_base)
-
- mean_dovish = out_dovish.mean().item()
- mean_hawkish = out_hawkish.mean().item()
-
- print(f" Dovish forecast mean : {mean_dovish:+.6f}")
- print(f" Hawkish forecast mean : {mean_hawkish:+.6f}")
- print(f" Δ (dovish − hawkish) : {mean_dovish - mean_hawkish:+.6f}")
-
- # The model's outputs must at least be numerically distinct for two
- # semantically different events (even in uncalibrated zero-shot mode).
- if torch.allclose(out_dovish, out_hawkish, atol=1e-6):
- _fail(
- "Dovish and hawkish forecasts are identical – "
- "text conditioning may not be working."
- )
- _ok("Dovish and hawkish forecasts are numerically distinct ✓")
-
- # ══════════════════════════════════════════════════════════════════════════
- # Summary
- # ══════════════════════════════════════════════════════════════════════════
- _sep("All 7 tests passed ✓")
- sys.exit(0)
diff --git a/singular_ticker_causal/services/pelt_detection.py b/singular_ticker_causal/services/pelt_detection.py
deleted file mode 100644
index d8f4cc34856e4257645b3487f28b6497a2cfb0a8..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/services/pelt_detection.py
+++ /dev/null
@@ -1,453 +0,0 @@
-import logging
-from typing import Dict, List, Optional
-
-import numpy as np
-import pandas as pd
-import ruptures as rpt
-from sklearn.preprocessing import StandardScaler
-
-
-logger = logging.getLogger(__name__)
-
-
-class PELTDetector:
- """
- Offline changepoint detection using the PELT algorithm
- from the ruptures library.
-
- Optimized for:
- - stock regime detection
- - structural break analysis
- - volatility shifts
- - factor regime segmentation
- """
-
- def __init__(
- self,
- model: str = "rbf",
- min_size: int = 20,
- jump: int = 2,
- penalty_scale: float = 3.0,
- use_log_returns: bool = True,
- ):
- """
- Parameters
- ----------
- model : str
- Cost model:
- - "l2" : mean shifts
- - "rbf" : nonlinear regime changes
- - "linear"
- - "normal"
- - "ar"
-
- min_size : int
- Minimum regime length.
-
- jump : int
- Subsampling factor for speed.
-
- penalty_scale : float
- Controls sensitivity.
- Higher = fewer changepoints.
-
- use_log_returns : bool
- Convert price series to log returns before detection.
- """
- self.model = model
- self.min_size = min_size
- self.jump = jump
- self.penalty_scale = penalty_scale
- self.use_log_returns = use_log_returns
-
- # ---------------------------------------------------------
- # Feature preprocessing
- # ---------------------------------------------------------
-
- def _prepare_series(
- self,
- series: pd.Series,
- ) -> tuple[np.ndarray, pd.Index]:
-
- s = (
- series.astype(float)
- .replace([np.inf, -np.inf], np.nan)
- .ffill()
- .bfill()
- )
-
- if self.use_log_returns:
- # clip before log to guard against zeros, negative values,
- # and bad corporate-action adjustments in real market data
- s = np.log(s.clip(lower=1e-8)).diff()
- s = s.replace([np.inf, -np.inf], np.nan).dropna()
-
- values = s.values.reshape(-1, 1)
-
- return values, s.index
-
- # ---------------------------------------------------------
- # Penalty estimation
- # ---------------------------------------------------------
-
- def _estimate_penalty(self, signal: np.ndarray) -> float:
- """
- Adaptive penalty estimation using MAD-based robust variance.
-
- MAD / 0.6745 is a consistent estimator of the standard deviation
- under normality and is substantially more robust to outliers and
- fat tails than the sample standard deviation — important for
- equity returns with volatility clustering.
- """
- flat = signal.ravel()
- mad = np.median(np.abs(flat - np.median(flat)))
- robust_variance = (mad / 0.6745) ** 2
- penalty = self.penalty_scale * np.log(len(signal)) * robust_variance
- return float(max(penalty, 1e-8))
-
- # ---------------------------------------------------------
- # Single-series detection
- # ---------------------------------------------------------
-
- def detect_changepoints(
- self,
- series: pd.Series,
- penalty: Optional[float] = None,
- ) -> List[pd.Timestamp]:
-
- try:
- signal, index = self._prepare_series(series)
-
- if len(signal) < self.min_size * 2:
- return []
-
- algo = rpt.Pelt(
- model=self.model,
- min_size=self.min_size,
- jump=self.jump,
- ).fit(signal)
-
- if penalty is None:
- penalty = self._estimate_penalty(signal)
-
- breakpoints = algo.predict(pen=penalty)
-
- # ruptures includes final endpoint — drop it
- breakpoints = breakpoints[:-1]
-
- dates = [
- index[min(bp, len(index) - 1)]
- for bp in breakpoints
- ]
-
- return dates
-
- except Exception as e:
- logger.exception(f"PELT changepoint detection failed: {e}")
- return []
-
- # ---------------------------------------------------------
- # Multivariate detection
- # ---------------------------------------------------------
-
- def detect_multivariate_changepoints(
- self,
- df: pd.DataFrame,
- columns: Optional[List[str]] = None,
- penalty: Optional[float] = None,
- ) -> List[pd.Timestamp]:
-
- try:
- if columns is None:
- columns = list(df.select_dtypes(include=np.number).columns)
-
- if not columns:
- return []
-
- # Build each column as a Series so NaN patterns are preserved
- # per-column. Stacking raw numpy arrays risks silent index
- # misalignment when columns have different NaN positions after
- # differencing.
- processed: Dict[str, pd.Series] = {}
-
- for col in columns:
- s = (
- df[col]
- .astype(float)
- .replace([np.inf, -np.inf], np.nan)
- .ffill()
- .bfill()
- )
-
- if self.use_log_returns:
- s = np.log(s.clip(lower=1e-8)).diff()
- s = s.replace([np.inf, -np.inf], np.nan)
-
- processed[col] = s
-
- # dropna across all columns simultaneously — ensures every row
- # is complete before we hand the matrix to PELT
- feature_df = pd.DataFrame(processed, index=df.index).dropna()
-
- if len(feature_df) < self.min_size * 2:
- return []
-
- # Standardize so rbf distances are not dominated by whichever
- # feature has the largest absolute scale (e.g. volume vs returns)
- scaler = StandardScaler()
- signal = scaler.fit_transform(feature_df.values)
-
- algo = rpt.Pelt(
- model=self.model,
- min_size=self.min_size,
- jump=self.jump,
- ).fit(signal)
-
- if penalty is None:
- penalty = self._estimate_penalty(signal)
-
- breakpoints = algo.predict(pen=penalty)
- breakpoints = breakpoints[:-1]
-
- dates = [
- feature_df.index[min(bp, len(feature_df.index) - 1)]
- for bp in breakpoints
- ]
-
- return dates
-
- except Exception as e:
- logger.exception(f"Multivariate PELT detection failed: {e}")
- return []
-
- # ---------------------------------------------------------
- # Column-wise detection
- # ---------------------------------------------------------
-
- def detect_dataframe_changepoints(
- self,
- df: pd.DataFrame,
- ) -> Dict[str, List[pd.Timestamp]]:
-
- results = {}
-
- numeric_cols = df.select_dtypes(include=np.number).columns
-
- for col in numeric_cols:
- try:
- results[col] = self.detect_changepoints(df[col])
- except Exception as e:
- logger.exception(f"Failed on column {col}: {e}")
- results[col] = []
-
- return results
-
- # ---------------------------------------------------------
- # Post-processing utility
- # ---------------------------------------------------------
-
- @staticmethod
- def merge_nearby_breakpoints(
- breakpoints: List[pd.Timestamp],
- index: pd.Index,
- min_gap: int = 10,
- ) -> List[pd.Timestamp]:
- """
- Merge changepoints that are closer than ``min_gap`` samples apart.
-
- Earnings spikes and gap-up/gap-down days can produce clusters of
- spurious micro-regime boundaries. This collapses each cluster to
- its first member.
-
- Parameters
- ----------
- breakpoints : list of pd.Timestamp
- Changepoint dates as returned by detect_changepoints or
- detect_multivariate_changepoints.
- index : pd.Index
- The DatetimeIndex of the series the breakpoints came from.
- Used to convert dates back to integer positions for gap
- measurement.
- min_gap : int
- Minimum number of samples that must separate two retained
- changepoints. Tune per use-case (e.g. 10 for daily NIFTY50,
- 30 for noisy smallcaps).
-
- Returns
- -------
- list of pd.Timestamp
- Filtered changepoints with near-duplicates removed.
- """
- if not breakpoints:
- return []
-
- positions = [index.get_loc(d) for d in breakpoints]
- merged_positions = [positions[0]]
-
- for pos in positions[1:]:
- if pos - merged_positions[-1] >= min_gap:
- merged_positions.append(pos)
-
- return [index[p] for p in merged_positions]
-
-
-def main():
- """
- Self-contained test for PELTDetector using synthetic data.
-
- Runs three scenarios:
- 1. Mean-shift series — detects breaks between segments of differing means.
- 2. Variance-change series — detects breaks between segments of differing
- variances (same mean), testing second-moment sensitivity.
- 3. Multi-column DataFrame — exercises detect_dataframe_changepoints on a
- DataFrame whose columns are drawn from both synthetic generators.
-
- Note: PELTDetector.use_log_returns=False is used here because the synthetic
- data is already stationary (zero-mean noise with added shifts); applying log
- returns would distort the signal.
- """
- import numpy as np
-
- logging.basicConfig(
- level=logging.INFO,
- format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
- )
- log = logging.getLogger("bocd_detector.main")
-
- # use_log_returns=False: synthetic data is already stationary
- detector = PELTDetector(
- model="rbf",
- min_size=20,
- jump=2,
- penalty_scale=3.0,
- use_log_returns=False,
- )
-
- # ------------------------------------------------------------------ #
- # Helper: build a DatetimeIndex of business days #
- # ------------------------------------------------------------------ #
- def make_date_index(n: int, start: str = "2020-01-02") -> pd.DatetimeIndex:
- return pd.bdate_range(start=start, periods=n)
-
- # ------------------------------------------------------------------ #
- # Synthetic data helpers (replaces generate_mean_shift_example etc.) #
- # ------------------------------------------------------------------ #
- rng = np.random.default_rng(0)
-
- def generate_mean_shift(
- num_segments: int = 4,
- segment_length: int = 80,
- shift_magnitude: float = 4.0,
- noise_std: float = 1.0,
- ) -> tuple[np.ndarray, np.ndarray]:
- """Returns (segment_lengths, data)."""
- segments = []
- for i in range(num_segments):
- mean = shift_magnitude if i % 2 == 1 else 0.0
- segments.append(rng.normal(mean, noise_std, segment_length))
- data = np.concatenate(segments)
- lengths = np.full(num_segments, segment_length)
- return lengths, data
-
- def generate_variance_change(
- num_segments: int = 3,
- segment_length: int = 120,
- variance_levels: Optional[List[float]] = None,
- ) -> tuple[np.ndarray, np.ndarray]:
- """Returns (segment_lengths, data)."""
- if variance_levels is None:
- variance_levels = [0.5, 4.0, 0.8]
- segments = []
- for var in variance_levels[:num_segments]:
- segments.append(rng.normal(0.0, np.sqrt(var), segment_length))
- data = np.concatenate(segments)
- lengths = np.full(num_segments, segment_length)
- return lengths, data
-
- # ------------------------------------------------------------------ #
- # Scenario 1 — Mean-shift series #
- # 4 segments × 80 pts, alternating mean 0 ↔ 4 #
- # Expected breaks near indices 80, 160, 240 #
- # ------------------------------------------------------------------ #
- log.info("=== Scenario 1: Mean-shift series ===")
- lengths_ms, data_ms = generate_mean_shift(
- num_segments=4, segment_length=80, shift_magnitude=4.0, noise_std=1.0
- )
- series_ms = pd.Series(
- data_ms,
- index=make_date_index(len(data_ms)),
- name="mean_shift",
- )
- true_breaks_ms = list(np.cumsum(lengths_ms)[:-1])
- log.info(f" Series length : {len(series_ms)}")
- log.info(f" True breaks : indices {true_breaks_ms}")
-
- cp_ms = detector.detect_changepoints(series_ms)
- log.info(f" Detected dates: {cp_ms}")
- if cp_ms:
- detected_idx = [series_ms.index.get_loc(d) for d in cp_ms]
- log.info(f" Detected idx : {detected_idx}")
- for idx in detected_idx:
- nearest = min(abs(idx - tb) for tb in true_breaks_ms)
- status = "✓" if nearest <= 15 else "✗ (far from true break)"
- log.info(f" idx={idx:3d} nearest_true={nearest:3d} {status}")
- else:
- log.warning(" No changepoints detected in mean-shift series.")
-
- # ------------------------------------------------------------------ #
- # Scenario 2 — Variance-change series #
- # 3 segments × 120 pts, variances [0.5, 4.0, 0.8], zero mean #
- # Expected breaks near indices 120, 240 #
- # ------------------------------------------------------------------ #
- log.info("=== Scenario 2: Variance-change series ===")
- lengths_vc, data_vc = generate_variance_change(
- num_segments=3,
- segment_length=120,
- variance_levels=[0.5, 4.0, 0.8],
- )
- series_vc = pd.Series(
- data_vc,
- index=make_date_index(len(data_vc)),
- name="variance_change",
- )
- true_breaks_vc = list(np.cumsum(lengths_vc)[:-1])
- log.info(f" Series length : {len(series_vc)}")
- log.info(f" True breaks : indices {true_breaks_vc}")
-
- cp_vc = detector.detect_changepoints(series_vc)
- log.info(f" Detected dates: {cp_vc}")
- if cp_vc:
- detected_idx = [series_vc.index.get_loc(d) for d in cp_vc]
- log.info(f" Detected idx : {detected_idx}")
- for idx in detected_idx:
- nearest = min(abs(idx - tb) for tb in true_breaks_vc)
- status = "✓" if nearest <= 15 else "✗ (far from true break)"
- log.info(f" idx={idx:3d} nearest_true={nearest:3d} {status}")
- else:
- log.warning(" No changepoints detected in variance-change series.")
-
- # ------------------------------------------------------------------ #
- # Scenario 3 — Multi-column DataFrame #
- # Two synthetic columns (mean-shift + variance-change) combined #
- # ------------------------------------------------------------------ #
- log.info("=== Scenario 3: Multi-column DataFrame ===")
- n_pts = min(len(series_ms), len(series_vc))
- df = pd.DataFrame(
- {
- "mean_shift": series_ms.values[:n_pts],
- "variance_change": series_vc.values[:n_pts],
- "price_proxy": series_ms.values[:n_pts] + 0.5 * series_vc.values[:n_pts],
- },
- index=make_date_index(n_pts),
- )
- log.info(f" DataFrame shape: {df.shape}")
-
- cp_df = detector.detect_dataframe_changepoints(df)
- for col, dates in cp_df.items():
- log.info(f" [{col}] → {len(dates)} changepoint(s) detected: {dates}")
-
- log.info("=== PELTDetector smoke-test complete ===")
-
-
-if __name__ == "__main__":
- main()
\ No newline at end of file
diff --git a/singular_ticker_causal/services/schema.py b/singular_ticker_causal/services/schema.py
deleted file mode 100644
index 5f5d986f6e7b1ee49a98d82cc9e739b749151c35..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/services/schema.py
+++ /dev/null
@@ -1,49 +0,0 @@
-# singular_ticker_causal/cuts_plus/data/schema.py
-
-# Layer 1 — Core DuPont & Income Statement
-INCOME_STATEMENT_NODES = [
- "Revenue",
- "COGS", # Cost of Materials Consumed
- "Operating_Expenses", # SG&A aggregate
- "EBIT", # Operating Income
- "Interest_Expense", # Finance Costs
- "EBT", # Earnings Before Tax
- "Tax_Expense",
- "PAT", # Net Income / Profit After Tax
- "Exceptional_Items", # Ind AS 1 / Schedule III mandate
- "OCI", # Other Comprehensive Income (Ind AS 1)
- "Total_Comprehensive_Income",# PAT + OCI — true driver of Total_Equity
-]
-
-# Layer 2 — Core Balance Sheet
-BALANCE_SHEET_NODES = [
- "Average_Total_Assets", # for Asset Turnover calculation
- "Average_Shareholders_Equity",# for Equity Multiplier calculation
- "PPE", # Property, Plant & Equipment (Ind AS 16)
- "Intangible_Assets", # Ind AS 38 — separate depreciation rules
- "ROU_Assets", # Ind AS 116 — Right-of-Use Assets
- "Lease_Liabilities", # Ind AS 116 — Lease Liability
- "Accounts_Receivable_Gross",
- "ECL_Allowance", # Ind AS 109 — Expected Credit Loss
- "Accounts_Receivable_Net", # = Gross - ECL_Allowance
- "Inventory",
- "Accounts_Payable",
- "Total_Debt", # includes Lease_Liabilities post Ind AS 116
- "CWIP", # Capital Work-in-Progress (Schedule III aging)
-]
-
-# Layer 3 — Strategic KPI / DuPont Outcome Nodes
-STRATEGIC_NODES = [
- "Gross_Profit", # Revenue - COGS
- "EBITDA", # EBIT + D&A (Note: post-116 rent excluded)
- "Net_Profit_Margin", # PAT / Revenue
- "Asset_Turnover", # Revenue / Average_Total_Assets
- "Equity_Multiplier", # Average_Total_Assets / Average_Shareholders_Equity
- "ROE", # Net_Profit_Margin × Asset_Turnover × Equity_Multiplier
- "ROCE",
- "Operating_Cash_Flow",
- "Capex",
- "Free_Cash_Flow", # OCF - Capex
- "Revolver_Borrowings", # plugless model: auto-draw on cash deficit
- "NCI", # Non-Controlling Interest (Ind AS 110)
-]
diff --git a/singular_ticker_causal/services/ssa_denoiser.py b/singular_ticker_causal/services/ssa_denoiser.py
deleted file mode 100644
index 9b0162254dc8dbc077ff7e6f3a1b81f270e7a596..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/services/ssa_denoiser.py
+++ /dev/null
@@ -1,738 +0,0 @@
-"""
-MSSA Quant Engine — Multivariate Singular Spectrum Analysis for financial signal extraction.
-
-Known limitations / future upgrades:
- - Issue 7 (VSSA): Currently uses Recurrent SSA forecasting. Vector SSA (VSSA) is generally
- more robust under structural breaks, nonlinear dynamics, and outlier-heavy data that are
- common in equity markets. Recurrent SSA is retained for its simplicity; VSSA should be
- evaluated as a future upgrade.
- - Issue 6 (Rolling MSSA): A rolling_fit() stub is provided but not yet implemented. This
- is essential for regime adaptation in nonstationary markets.
- - Issue 2 (Eigentriple grouping): Oscillatory financial modes appear as paired eigentriples
- with nearly equal singular values. True grouping (by pairing + w-correlation clustering)
- is the next major upgrade beyond the current top-k selection.
-"""
-
-import numpy as np
-import pandas as pd
-import logging
-from scipy.sparse.linalg import svds
-from sklearn.utils.extmath import randomized_svd
-from sklearn.preprocessing import StandardScaler
-from sklearn.covariance import LedoitWolf
-from sklearn.mixture import GaussianMixture
-import yfinance as yf
-
-
-logger = logging.getLogger(__name__)
-
-
-class MSSAQuantEngine:
- """
- Multivariate Singular Spectrum Analysis (MSSA) framework for signal extraction,
- forecasting, regime detection, and covariance denoising.
- """
-
- def __init__(
- self,
- window_size=60,
- rank=None,
- variance_threshold=0.90,
- use_randomized_svd=True,
- random_state=42
- ):
- self.window_size = window_size
- self.rank = rank
- self.variance_threshold = variance_threshold
- self.use_randomized_svd = use_randomized_svd
- self.random_state = random_state
-
- self.U = None
- self.S = None
- self.VT = None
-
- self.components_ = None
- self.grouped_components_ = None
-
- # =========================================================
- # DATA INGESTION
- # =========================================================
-
- @staticmethod
- def download_indian_data(
- tickers,
- start="2015-01-01",
- end=None,
- interval="1d"
- ):
- yf_tickers = [f"{t}.NS" for t in tickers]
-
- data = yf.download(
- yf_tickers,
- start=start,
- end=end,
- interval=interval,
- auto_adjust=True,
- progress=False
- )
-
- close = data["Close"]
- close = close.dropna(how="all")
-
- return close
-
- # =========================================================
- # FEATURE ENGINEERING
- # =========================================================
-
- @staticmethod
- def compute_features(price_df):
- returns = np.log(price_df / price_df.shift(1))
-
- realized_vol = (
- returns.rolling(20).std() * np.sqrt(252)
- )
-
- momentum = returns.rolling(20).mean()
-
- drawdown = (
- price_df / price_df.rolling(252).max() - 1
- )
-
- features = pd.concat(
- {
- "returns": returns,
- "volatility": realized_vol,
- "momentum": momentum,
- "drawdown": drawdown,
- },
- axis=1
- )
-
- return features.dropna()
-
- # =========================================================
- # MSSA EMBEDDING
- # =========================================================
-
- def _trajectory_matrix(self, series):
- x = series.values
- N = len(x)
- L = self.window_size
- K = N - L + 1
-
- return np.column_stack(
- [x[i:i + L] for i in range(K)]
- )
-
- def build_block_hankel(self, df):
- matrices = []
- for col in df.columns:
- s = df[col].ffill().bfill()
- X = self._trajectory_matrix(s)
- matrices.append(X)
-
- return np.vstack(matrices)
-
- # =========================================================
- # SVD
- # =========================================================
-
- def fit(self, df, train_end=None):
- """
- Fit the MSSA engine on *df*.
-
- Parameters
- ----------
- df : pd.DataFrame
- Input feature matrix (time × assets/features).
- train_end : int or None
- If provided, the StandardScaler is fitted **only on the first
- `train_end` rows** and then applied to the full DataFrame. This
- prevents future-data leakage into the normalisation step, which
- would produce overly optimistic backtests.
- When None (default), the scaler is fitted on the whole DataFrame
- (acceptable for research / exploratory use, but not for live
- trading or walk-forward validation).
- """
- self.columns_ = df.columns
- self.index_ = df.index
-
- scaler = StandardScaler()
-
- # fit scaler only on the designated training window.
- if train_end is not None:
- scaler.fit(df.iloc[:train_end])
- else:
- logger.warning(
- "fit() called without train_end — scaler is fitted on the full "
- "dataset, which leaks future information into normalisation. "
- "Pass train_end= for walk-forward safe usage."
- )
- scaler.fit(df)
-
- scaled = pd.DataFrame(
- scaler.transform(df),
- columns=df.columns,
- index=df.index
- )
-
- self.scaler_ = scaler
- X = self.build_block_hankel(scaled)
- self.X_ = X
-
- if self.use_randomized_svd:
- rank = min(50, min(X.shape) - 1)
- U, S, VT = randomized_svd(
- X,
- n_components=rank,
- random_state=self.random_state
- )
- else:
- rank = min(50, min(X.shape) - 1)
- U, S, VT = svds(X, k=rank)
-
- idx = np.argsort(S)[::-1]
- S = S[idx]
- U = U[:, idx]
- VT = VT[idx]
-
- self.U = U
- self.S = S
- self.VT = VT
-
- self.rank_ = self._automatic_rank_selection()
-
- return self
-
- # =========================================================
- # AUTOMATIC RANK SELECTION
- # =========================================================
-
- def _automatic_rank_selection(self):
- if self.rank is not None:
- return self.rank
-
- eigvals = self.S ** 2
- explained = eigvals / eigvals.sum()
- cumulative = np.cumsum(explained)
-
- rank = (
- np.searchsorted(
- cumulative,
- self.variance_threshold
- ) + 1
- )
-
- return rank
-
- # =========================================================
- # COMPONENT GROUPING
- # =========================================================
-
- def reconstruct_group(self, group):
- X_recon = np.zeros_like(self.X_)
-
- for i in group:
- X_recon += (
- self.S[i] * np.outer(self.U[:, i], self.VT[i])
- )
-
- return X_recon
-
- def _diagonal_averaging(self, X_block):
- """
- Converts block Hankel matrix back into multivariate time series.
- Builds a plain numpy array first to avoid issues with duplicate
- column names in intermediate DataFrames.
- """
-
- L = self.window_size
- K = X_block.shape[1]
- N = L + K - 1
- n_series = len(self.columns_)
-
- # Build (N, n_series) numpy array via diagonal averaging
- result = np.zeros((N, n_series))
-
- for idx in range(n_series):
- start = idx * L
- end = (idx + 1) * L
- X = X_block[start:end] # shape (L, K)
-
- ts = np.zeros(N)
- counts = np.zeros(N)
-
- for i in range(L):
- for j in range(K):
- ts[i + j] += X[i, j]
- counts[i + j] += 1
-
- ts /= counts
- result[:, idx] = ts
-
- return pd.DataFrame(
- result,
- index=self.index_[:N],
- columns=self.columns_
- )
-
- def extract_signal(self):
- """
- Extract the reconstructed signal using the top-k eigentriples.
-
- Issue 2 (Eigentriple Grouping) — TODO:
- The current grouping strategy (top ``rank_`` eigentriples) assumes
- that the leading singular values always correspond to the trend /
- signal components. In practice, financial oscillatory modes appear
- as *paired* eigentriples with nearly equal singular values. A more
- principled grouping based on eigenvalue pairing, phase similarity,
- and w-correlation clustering should replace this simple top-k
- selection in a future upgrade.
- """
-
- signal_group = list(range(self.rank_))
-
- X_signal = self.reconstruct_group(signal_group)
-
- reconstructed = self._diagonal_averaging(X_signal)
-
- # Pass .values so sklearn sees a plain numpy array,
- # avoiding any duplicate-column confusion.
- inverse = self.scaler_.inverse_transform(reconstructed.values)
-
- return pd.DataFrame(
- inverse,
- columns=reconstructed.columns,
- index=reconstructed.index
- )
-
- def w_correlation(self, reconstructed_df):
- """
- Compute pairwise correlation between reconstructed components.
-
- Issue 1 (W-Correlation) — Current status and TODO:
- This method currently computes **ordinary Pearson correlation**
- (``np.corrcoef``), which is a practical approximation.
-
- Canonical SSA w-correlation uses *weighted* inner products that
- respect the Hankel structure of the trajectory matrix. The weight
- for lag *k* is:
-
- w(k) = min(k+1, L, N-L+1, N-k) (L = window size, N = series length)
-
- TODO: Replace ``np.corrcoef`` with a ``_weighted_inner_product()``
- helper that applies these diagonal weights. This is needed for
- theoretically correct eigentriple grouping.
- """
-
- X = reconstructed_df.values
-
- n = X.shape[1]
-
- wcorr = np.zeros((n, n))
-
- for i in range(n):
- for j in range(n):
-
- xi = X[:, i]
- xj = X[:, j]
-
- wcorr[i, j] = np.corrcoef(xi, xj)[0, 1]
-
- return wcorr
-
- # =========================================================
- # FORECASTING
- # =========================================================
-
- def recurrent_forecast(self, series, steps=5):
- """
- Recurrent SSA forecast for a single time series.
-
- Forecast Horizon Stability:
- Recurrent SSA forecasts accumulate approximation error and become
- increasingly unstable at long horizons. Empirical guidance:
-
- ========= =================
- Horizon Stability
- ========= =================
- 1 – 5 Good
- 5 – 20 Acceptable
- 20+ Unstable / unreliable
- ========= =================
-
- A warning is raised when ``steps > 5``.
-
- Issue 7 (VSSA) — TODO:
- Vector SSA (VSSA) is generally more robust than Recurrent SSA
- for nonstationary and outlier-heavy financial series. Consider
- implementing VSSA as an alternative forecasting backend.
- """
-
- # warn on long-horizon forecasts.
- if steps > 5:
- logger.warning(
- "recurrent_forecast called with steps=%d. Recurrent SSA "
- "forecasts are known to become unstable beyond 5 steps. "
- "Use short-horizon forecasts (steps <= 5) for reliable results.",
- steps,
- )
-
- s = series.values.copy()
- L = self.window_size
-
- X = self._trajectory_matrix(pd.Series(s))
-
- U, S, VT = np.linalg.svd(
- X,
- full_matrices=False
- )
-
- r = min(self.rank_, len(S))
- Ur = U[:, :r]
-
- pi = Ur[-1]
- nu = Ur[:-1]
-
- coeffs = nu @ pi / (1 - np.sum(pi ** 2))
-
- forecasts = []
- extended = list(s)
-
- for _ in range(steps):
- next_value = np.dot(
- coeffs[::-1],
- extended[-len(coeffs):]
- )
- forecasts.append(next_value)
- extended.append(next_value)
-
- return forecasts
-
- def forecast_signal(self, reconstructed_df, steps=5):
- forecasts = {}
- for col_idx in range(len(reconstructed_df.columns)):
- col_name = reconstructed_df.columns[col_idx]
- series = reconstructed_df.iloc[:, col_idx]
- forecasts[f"{col_name}_{col_idx}"] = self.recurrent_forecast(series, steps=steps)
-
- df = pd.DataFrame(forecasts)
- df.columns = reconstructed_df.columns
- return df
-
- # =========================================================
- # REGIME FEATURES
- # =========================================================
-
- def regime_features(self):
- singular_ratio = (
- self.S[0] / self.S.sum()
- )
-
- spectral_entropy = -np.sum(
- (self.S / self.S.sum())
- * np.log(self.S / self.S.sum())
- )
-
- trend_strength = (
- np.sum(self.S[:self.rank_])
- / np.sum(self.S)
- )
-
- return {
- "singular_ratio": singular_ratio,
- "spectral_entropy": spectral_entropy,
- "trend_strength": trend_strength
- }
-
- # =========================================================
- # REGIME CLUSTERING
- # =========================================================
-
- @staticmethod
- def _label_regime(centroids):
- """
- Regime Labels Are Arbitrary.
-
- GMM cluster indices are arbitrary — the same data can yield a different
- index ordering on each run. This method maps each cluster to a
- semantically meaningful label by inspecting its centroid in the
- (singular_ratio, spectral_entropy, trend_strength) feature space:
-
- +---------------------------+------------------+
- | Feature signal | Regime |
- +===========================+==================+
- | High singular_ratio | TREND |
- | Low entropy | |
- +---------------------------+------------------+
- | High entropy | HIGH_VOLATILITY |
- | Low singular_ratio | |
- +---------------------------+------------------+
- | Medium entropy + | MEAN_REVERSION |
- | medium singular_ratio | |
- +---------------------------+------------------+
-
- Parameters
- ----------
- centroids : np.ndarray, shape (n_clusters, 3)
- GMM means in order [singular_ratio, spectral_entropy, trend_strength].
-
- Returns
- -------
- dict[int, str]
- Mapping from GMM cluster index → regime label.
- """
- labels = {}
- # Score each centroid on three criteria:
- # trending — high singular_ratio, low entropy
- # high_vol — high entropy, low singular_ratio
- # mean_rev — everything else (medium values)
- trend_scores = centroids[:, 0] - centroids[:, 1] # singular_ratio – entropy
- hv_scores = centroids[:, 1] - centroids[:, 0] # entropy – singular_ratio
-
- assigned = set()
- # Assign TREND to the cluster with highest (singular_ratio – entropy)
- trend_idx = int(np.argmax(trend_scores))
- labels[trend_idx] = "TREND"
- assigned.add(trend_idx)
-
- # Assign HIGH_VOLATILITY to the cluster with highest (entropy – singular_ratio)
- # among the remaining clusters
- masked_hv = hv_scores.copy()
- for idx in assigned:
- masked_hv[idx] = -np.inf
- hv_idx = int(np.argmax(masked_hv))
- labels[hv_idx] = "HIGH_VOLATILITY"
- assigned.add(hv_idx)
-
- # Remaining clusters get MEAN_REVERSION
- for idx in range(len(centroids)):
- if idx not in assigned:
- labels[idx] = "MEAN_REVERSION"
-
- return labels
-
- def detect_regime(self):
- """
- Detect the current market regime using a GMM fitted on historical
- trajectory-matrix sub-windows.
-
- Regime Labels Are Arbitrary fix:
- Regime labels are now assigned dynamically by ``_label_regime()``
- based on each cluster's centroid in feature space, rather than
- relying on the arbitrary integer ordering that GMM produces.
- """
- feats = self.regime_features()
-
- X_current = np.array([
- [
- feats["singular_ratio"],
- feats["spectral_entropy"],
- feats["trend_strength"]
- ]
- ])
-
- # Generate real historical data by performing SVD on rolling sub-matrices
- # of the trajectory matrix X_ to build a history of regime features.
- _, K = self.X_.shape
- sub_W = min(K // 2, 120)
-
- historical_X = []
- if sub_W >= 10:
- # Generate about 30 historical samples to fit the GMM
- step = max(1, (K - sub_W) // 30)
-
- for start in range(0, K - sub_W, step):
- X_sub = self.X_[:, start:start + sub_W]
-
- # Compute SVD for the sub-matrix
- rank_sub = min(30, min(X_sub.shape) - 1)
- try:
- if self.use_randomized_svd:
- _, S_sub, _ = randomized_svd(
- X_sub, n_components=rank_sub, random_state=self.random_state
- )
- else:
- _, S_sub, _ = svds(X_sub, k=rank_sub)
- S_sub = np.sort(S_sub)[::-1]
-
- S_sum = np.sum(S_sub)
- if S_sum == 0:
- continue
-
- sing_ratio = S_sub[0] / S_sum
- entropy = -np.sum((S_sub / S_sum) * np.log((S_sub / S_sum) + 1e-9))
-
- eigvals = S_sub ** 2
- explained = eigvals / eigvals.sum()
- cumulative = np.cumsum(explained)
- sub_rank_val = np.searchsorted(cumulative, self.variance_threshold) + 1
- trend_str = np.sum(S_sub[:sub_rank_val]) / S_sum
-
- historical_X.append([sing_ratio, entropy, trend_str])
- except Exception:
- continue
-
- if len(historical_X) >= 3:
- X_train = np.vstack([historical_X, X_current])
- else:
- # Fallback if trajectory matrix is too small to generate history
- fallback_X = np.array([
- [0.8, 1.0, 0.9],
- [0.2, 3.0, 0.2],
- [0.5, 2.0, 0.5],
- ])
- X_train = np.vstack([fallback_X, X_current])
-
- gmm = GaussianMixture(
- n_components=min(3, len(X_train)),
- random_state=42
- )
-
- gmm.fit(X_train)
- regime_idx = int(gmm.predict(X_current)[0])
-
- # derive labels from centroid semantics, not arbitrary index.
- regime_labels = self._label_regime(gmm.means_)
- regime_label = regime_labels.get(regime_idx, "UNKNOWN")
-
- return regime_idx, regime_label
-
- # =========================================================
- # COVARIANCE DENOISING
- # =========================================================
-
- @staticmethod
- def denoise_covariance(returns_df):
- lw = LedoitWolf()
- lw.fit(returns_df.dropna())
- cov = lw.covariance_
-
- return pd.DataFrame(
- cov,
- index=returns_df.columns,
- columns=returns_df.columns
- )
-
- # =========================================================
- # ROLLING MSSA (Issue 6 — stub, not yet implemented)
- # =========================================================
-
- def rolling_fit(self, df, window=252, step=1):
- """
- Issue 6 (No Rolling MSSA) — TODO: implement.
-
- Markets are nonstationary. A single static fit is insufficient for
- live trading or regime-adaptive strategies. This method should:
-
- 1. Slide a window of length ``window`` rows over ``df``.
- 2. Call ``self.fit(df.iloc[t:t+window], train_end=window)`` at each
- step ``t``.
- 3. Collect the extracted signals and detected regimes at each step.
- 4. Return a time-indexed DataFrame of rolling signals and labels.
-
- This is essential for:
- * Regime adaptation to changing covariance structures
- * Walk-forward backtesting without look-ahead bias
- * Handling Indian macro structural breaks
-
- Parameters
- ----------
- df : pd.DataFrame
- Full input feature matrix.
- window : int
- Rolling window length in rows (default: 252 trading days).
- step : int
- Number of rows to advance the window at each iteration.
-
- Raises
- ------
- NotImplementedError
- Always — method is not yet implemented.
- """
- raise NotImplementedError(
- "rolling_fit() is not yet implemented. See the docstring for "
- "the planned implementation specification (Issue 6)."
- )
-
-
-def main():
- print("\n" + "=" * 60)
- print("TEST – MSSAQuantEngine Example Usage")
- print("=" * 60)
-
- try:
- tickers = [
- "RELIANCE",
- "TCS",
- "INFY",
- "HDFCBANK",
- "ICICIBANK",
- "SBIN",
- ]
-
- print(f"Downloading data for {tickers}...")
- prices = MSSAQuantEngine.download_indian_data(
- tickers,
- start="2018-01-01"
- )
-
- if prices is not None and not prices.empty:
- print("Computing features...")
- features = MSSAQuantEngine.compute_features(prices)
-
- # Rename columns to be globally unique before concatenating,
- # so that duplicate ticker names don't cause downstream issues.
- ret = features["returns"].add_suffix("_ret")
- vol = features["volatility"].add_suffix("_vol")
- mom = features["momentum"].add_suffix("_mom")
-
- feature_matrix = pd.concat([ret, vol, mom], axis=1)
-
- print("Fitting MSSA...")
- engine = MSSAQuantEngine(
- window_size=60,
- variance_threshold=0.90
- )
- # Prevent future data leakage by fitting the scaler on the first 80% of data
- train_split = int(len(feature_matrix) * 0.8)
- engine.fit(feature_matrix, train_end=train_split)
-
- print("Extracting signal...")
- signal = engine.extract_signal()
- print(f"Signal extracted. Shape: {signal.shape}")
-
- print("Forecasting signals...")
- forecast = engine.forecast_signal(signal, steps=5)
- print(f"Forecasted signals shape: {forecast.shape}")
- if "RELIANCE.NS" in forecast.columns:
- rel_forecast = forecast["RELIANCE.NS"]
- if isinstance(rel_forecast, pd.DataFrame):
- rel_forecast = rel_forecast.iloc[:, 0]
- print(f"Forecast (5 steps) for RELIANCE.NS SIGNAL:\n{rel_forecast.values}")
-
- print("Detecting regime...")
- # detect_regime() now returns (index, label) tuple;
- # labels are derived from centroid semantics, not arbitrary int mapping.
- regime_idx, regime_label = engine.detect_regime()
- print(f"Regime: {regime_idx} ({regime_label})")
-
- print("Denoising covariance...")
- cov = engine.denoise_covariance(features["returns"])
- print(f"Covariance shape: {cov.shape}")
-
- print("\n✓ Example usage completed successfully.")
- else:
- print("No data downloaded. Skipping further tests.")
-
- except Exception as e:
- logger.error(f"Error during example usage: {e}")
-
-
-if __name__ == "__main__":
- logging.basicConfig(
- level=logging.INFO,
- format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
- )
- main()
diff --git a/singular_ticker_causal/services/tensor_builder.py b/singular_ticker_causal/services/tensor_builder.py
deleted file mode 100644
index ea9a6e998a19d8c3c7d5a9bcf8ed9d7182d27f84..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/services/tensor_builder.py
+++ /dev/null
@@ -1,979 +0,0 @@
-import json
-import pandas as pd
-import logging
-import os
-import numpy as np
-import torch
-from transformers import BertTokenizer, BertModel
-from typing import List, Dict, Any, Optional, Tuple
-from datetime import datetime
-from .schema import (
- INCOME_STATEMENT_NODES,
- BALANCE_SHEET_NODES,
- STRATEGIC_NODES,
-)
-
-logger = logging.getLogger(__name__)
-
-DEBUG_DIR = os.path.join(os.path.dirname(__file__), "..", "debug_data")
-os.makedirs(DEBUG_DIR, exist_ok=True)
-
-
-# ── Node taxonomy ───────────────────────────────────────────────────────────
-# Discovery uses ONLY primitive (Income Statement + Balance Sheet) nodes.
-# Derived / KPI nodes (STRATEGIC_NODES) are algebraic identities and will
-# create fake causal edges if included in the discovery graph.
-# They must be computed AFTER graph inference via FinancialMetricEngine.
-DISCOVERY_NODES: List[str] = INCOME_STATEMENT_NODES + BALANCE_SHEET_NODES
-DERIVED_NODES: List[str] = STRATEGIC_NODES # used post-inference only
-
-# Per-modality exponential decay rates (λ). Larger λ → faster decay.
-# news decays quickly; policy persists ~2 years; governance in between.
-_DECAY_LAMBDA = {"news": 0.2, "policy": 0.03, "gov": 0.08}
-
-# FIX #10 — Explicit per-modality maximum horizon caps (months).
-# Acts as a hard ceiling on top of the 1e-4 early-stop threshold.
-_MAX_EVENT_HORIZON: Dict[str, int] = {"news": 3, "policy": 24, "gov": 12}
-
-# FIX #9 — Event tensor dimensionalities.
-# The sparse denoiser now emits low-dimensional causal event vectors
-# over the fundamental taxonomy, with separate positive/negative slots.
-_FUNDAMENTAL_TAXONOMY = (
- "Revenue, COGS, GrossProfit, EBITDA, EBIT, NetIncome, "
- "TotalAssets, TotalDebt, CashAndEquivalents, "
- "OperatingCashFlow, CapEx, FreeCashFlow, "
- "ShareholdersEquity, RetainedEarnings"
-)
-_EVENT_DIM = len(_FUNDAMENTAL_TAXONOMY.split(",")) * 2 # node×direction.
-_COMBINED_DIM = _EVENT_DIM
-_TEXT_DIM = _COMBINED_DIM * 3 # news + policy + gov
-
-# FIX #7 — Financial DAG priors: allowed causal directions between primitive nodes.
-# Used to build the adjacency mask passed to CUTS+. Only edges listed here
-# (src → tgt) are permitted; the rest are zeroed out.
-ALLOWED_CAUSAL_DIRECTIONS: Dict[str, List[str]] = {
- # ── Income Statement flow (top-down) ────────────────────────────────────
- # Revenue drives cost structures and receivables
- "Revenue": [
- "COGS",
- "Operating_Expenses",
- "Accounts_Receivable_Gross",
- "Inventory",
- "PAT",
- ],
-
- # COGS is driven by inventory consumed and payables incurred
- "COGS": [
- "EBIT",
- "Inventory",
- "Accounts_Payable",
- ],
-
- # Operating expenses reduce EBIT
- "Operating_Expenses": [
- "EBIT",
- ],
-
- # Depreciation feeds into EBIT (reduces it) and erodes PPE/ROU
- "Depreciation": [
- "EBIT",
- "PPE",
- "ROU_Assets",
- ],
-
- # EBIT minus interest gives EBT
- "EBIT": [
- "EBT",
- "Interest_Expense", # higher debt burden shows up as interest drag
- ],
-
- # Debt service cost flows to EBT
- "Interest_Expense": [
- "EBT",
- ],
-
- # EBT minus tax gives PAT
- "EBT": [
- "PAT",
- "Tax_Expense",
- ],
-
- # Tax is a function of EBT
- "Tax_Expense": [
- "PAT",
- ],
-
- # Exceptional items distort PAT and flow into Total Comprehensive Income
- "Exceptional_Items": [
- "PAT",
- "Total_Comprehensive_Income",
- ],
-
- # OCI (Other Comprehensive Income) feeds Total Comprehensive Income
- "OCI": [
- "Total_Comprehensive_Income",
- ],
-
- # PAT feeds Total Comprehensive Income and retained earnings proxy
- "PAT": [
- "Total_Comprehensive_Income",
- "Average_Shareholders_Equity", # retained earnings build equity
- ],
-
- # Total Comprehensive Income is the terminal P&L node — no outgoing
- # causal edges to other primitives (it's a summary line)
-
- # ── Balance Sheet — Assets ───────────────────────────────────────────────
- # Capex drives PPE and CWIP (capital work in progress)
- "Capex": [
- "PPE",
- "CWIP",
- "Depreciation", # more assets → higher future depreciation
- ],
-
- # CWIP converts to PPE when projects are commissioned
- "CWIP": [
- "PPE",
- ],
-
- # PPE drives depreciation charges
- "PPE": [
- "Depreciation",
- "Average_Total_Assets",
- ],
-
- # Intangibles affect average total assets
- "Intangible_Assets": [
- "Average_Total_Assets",
- ],
-
- # ROU Assets (lease right-of-use) drive depreciation and lease liabilities
- "ROU_Assets": [
- "Depreciation",
- "Lease_Liabilities",
- "Average_Total_Assets",
- ],
-
- # Inventory build affects COGS timing and payables
- "Inventory": [
- "COGS",
- "Accounts_Payable",
- "Average_Total_Assets",
- ],
-
- # Gross receivables less ECL gives net receivables
- "Accounts_Receivable_Gross": [
- "Accounts_Receivable_Net",
- "ECL_Allowance", # higher gross AR → higher expected credit loss
- ],
-
- # ECL provision reduces net receivables and hits P&L (operating expenses)
- "ECL_Allowance": [
- "Accounts_Receivable_Net",
- "Operating_Expenses", # bad debt charge flows through opex
- ],
-
- # Net receivables affect average total assets
- "Accounts_Receivable_Net": [
- "Average_Total_Assets",
- ],
-
- # ── Balance Sheet — Liabilities ──────────────────────────────────────────
- # Debt drives interest expense and affects average total assets (via leverage)
- "Total_Debt": [
- "Interest_Expense",
- "Average_Total_Assets", # debt-funded assets inflate the asset base
- "Average_Shareholders_Equity", # leverage dilutes equity ratios
- ],
-
- # Accounts payable is a liability funded by inventory purchases
- "Accounts_Payable": [
- "Average_Total_Assets", # working capital affects asset base
- ],
-
- # Lease liabilities drive interest-equivalent charges (IFRS 16 finance cost)
- "Lease_Liabilities": [
- "Interest_Expense",
- "Average_Total_Assets",
- ],
-
- # ── Equity / Asset aggregates ────────────────────────────────────────────
- # Average total assets and equity are downstream aggregates;
- # they have no outgoing primitive causal edges within this node set.
-}
-
-# FIX #8 — Regime feature names.
-REGIME_FEATURES: List[str] = [
- "bull_market",
- "bear_market",
- "high_inflation",
- "tightening_cycle",
- "earnings_expansion",
- "earnings_contraction",
- "high_volatility",
-]
-
-
-# ── EmbeddingService ────────────────────────────────────────────────────────
-
-class EmbeddingService:
- """
- Converts raw news articles into dual embeddings:
- • 768-D FinBERT sentiment vector (Vansh180/FinBERT-India-v1)
- • 1024-D bge-large semantic/causal vector (BAAI/bge-large-en-v1.5)
-
- Concatenated → 1792-D per article.
-
- FIX #9: Single FinBERT was weak for causal event semantics. Dual embedding
- keeps financial sentiment while adding a model trained for semantic similarity
- and causal retrieval tasks.
- """
-
- def __init__(
- self,
- sentiment_model_id: str = "Vansh180/FinBERT-India-v1",
- semantic_model_id: str = "BAAI/bge-base-en-v1.5",
- ):
- logger.info(
- f"Initializing EmbeddingService: sentiment={sentiment_model_id}, "
- f"semantic={semantic_model_id}"
- )
- self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
-
- # ── Sentiment model (FinBERT) ──────────────────────────────────────
- self.sent_tokenizer = BertTokenizer.from_pretrained(sentiment_model_id)
- self.sent_model = BertModel.from_pretrained(sentiment_model_id)
- self.sent_model.to(self.device)
- self.sent_model.eval()
-
- # ── Semantic/causal model (bge-large) ─────────────────────────────
- # Use sentence-transformers if available; fall back to bare HuggingFace.
- try:
- from sentence_transformers import SentenceTransformer
- self.sem_model = SentenceTransformer(semantic_model_id, device=str(self.device))
- self._use_st = True
- except ImportError:
- logger.warning(
- "sentence-transformers not installed. Falling back to HuggingFace for "
- "semantic model. Install with: pip install sentence-transformers"
- )
- from transformers import AutoTokenizer, AutoModel
- self.sem_tokenizer = AutoTokenizer.from_pretrained(semantic_model_id)
- self.sem_hf_model = AutoModel.from_pretrained(semantic_model_id)
- self.sem_hf_model.to(self.device)
- self.sem_hf_model.eval()
- self._use_st = False
-
- # ── Internal helpers ───────────────────────────────────────────────────
-
- @staticmethod
- def _mean_pooling(model_output, attention_mask) -> torch.Tensor:
- token_embeddings = model_output.last_hidden_state
- input_mask_expanded = (
- attention_mask.unsqueeze(-1).expand(token_embeddings.size()).float()
- )
- return torch.sum(token_embeddings * input_mask_expanded, 1) / torch.clamp(
- input_mask_expanded.sum(1), min=1e-9
- )
-
- def _sentiment_embed_batch(self, texts: List[str]) -> np.ndarray:
- inputs = self.sent_tokenizer(
- texts, return_tensors="pt", padding=True, truncation=True, max_length=512
- )
- inputs = {k: v.to(self.device) for k, v in inputs.items()}
- with torch.no_grad():
- outputs = self.sent_model(**inputs)
- embs = self._mean_pooling(outputs, inputs["attention_mask"])
- return embs.cpu().numpy()
-
- def _semantic_embed_batch(self, texts: List[str]) -> np.ndarray:
- if self._use_st:
- return self.sem_model.encode(
- texts, batch_size=32, normalize_embeddings=True, show_progress_bar=False
- )
- # Fallback: bare HuggingFace mean-pool
- inputs = self.sem_tokenizer(
- texts, return_tensors="pt", padding=True, truncation=True, max_length=512
- )
- inputs = {k: v.to(self.device) for k, v in inputs.items()}
- with torch.no_grad():
- outputs = self.sem_hf_model(**inputs)
- embs = self._mean_pooling(outputs, inputs["attention_mask"])
- return embs.cpu().numpy()
-
- # ── Public API ─────────────────────────────────────────────────────────
-
- def embed_text(self, text: str) -> np.ndarray:
- """Embed a single string → 1792-D vector [sentiment || semantic]."""
- sent_emb = self._sentiment_embed_batch([text])[0] # (768,)
- sem_emb = self._semantic_embed_batch([text])[0] # (1024,)
- return np.concatenate([sent_emb, sem_emb]) # (1792,)
-
- def embed_news(
- self, news_by_symbol: Dict[str, List[dict]]
- ) -> Dict[str, List[dict]]:
- """
- Embeds all articles for all symbols in batched mode.
-
- Returns dict mapping symbol → list of article dicts with
- "embedding" key (1792-D list) added.
- """
- result: Dict[str, List[dict]] = {}
- all_texts: List[str] = []
- mapping: List[Tuple[str, dict]] = [] # (symbol, article_dict)
-
- for sym, articles in news_by_symbol.items():
- if not articles:
- result[sym] = []
- continue
- for article in articles:
- title = article.get("title", "").strip()
- desc = article.get("description", article.get("content", "")).strip()
- text = f"{title}. {desc}" if desc and desc != title else title
- all_texts.append(text)
- mapping.append((sym, article))
-
- if not all_texts:
- return result
-
- logger.info(
- f"Embedding {len(all_texts)} articles across {len(news_by_symbol)} symbols "
- f"(dual model: {_SENTIMENT_DIM}-D + {_SEMANTIC_DIM}-D = {_COMBINED_DIM}-D each)..."
- )
-
- batch_size = 32
- sent_embs_all: List[np.ndarray] = []
- sem_embs_all: List[np.ndarray] = []
-
- for i in range(0, len(all_texts), batch_size):
- batch = all_texts[i : i + batch_size]
- sent_embs_all.append(self._sentiment_embed_batch(batch))
- sem_embs_all.append(self._semantic_embed_batch(batch))
-
- sent_embs = np.vstack(sent_embs_all) # (n, 768)
- sem_embs = np.vstack(sem_embs_all) # (n, 1024)
- all_embeddings = np.concatenate([sent_embs, sem_embs], axis=1) # (n, 1792)
-
- for i, (sym, article) in enumerate(mapping):
- result.setdefault(sym, [])
- enriched = article.copy()
- enriched["embedding"] = all_embeddings[i].tolist()
- result[sym].append(enriched)
-
- return result
-
-
-# ── FinancialMetricEngine ────────────────────────────────────────────────────
-
-class FinancialMetricEngine:
- """
- Post-inference layer that computes deterministic KPI / derived nodes
- from the primitives discovered by the causal graph.
-
- These should NEVER be included in the causal discovery graph directly,
- as they are algebraic identities that produce fake causal edges.
- """
-
- @staticmethod
- def compute_gross_profit(revenue: np.ndarray, cogs: np.ndarray) -> np.ndarray:
- return revenue - cogs
-
- @staticmethod
- def compute_ebitda(ebit: np.ndarray, da: np.ndarray) -> np.ndarray:
- return ebit + da
-
- @staticmethod
- def compute_net_profit_margin(pat: np.ndarray, revenue: np.ndarray) -> np.ndarray:
- return np.where(revenue != 0, pat / revenue, 0.0)
-
- @staticmethod
- def compute_asset_turnover(revenue: np.ndarray, avg_assets: np.ndarray) -> np.ndarray:
- return np.where(avg_assets != 0, revenue / avg_assets, 0.0)
-
- @staticmethod
- def compute_equity_multiplier(avg_assets: np.ndarray, avg_equity: np.ndarray) -> np.ndarray:
- return np.where(avg_equity != 0, avg_assets / avg_equity, 0.0)
-
- @staticmethod
- def compute_roe(
- pat: np.ndarray,
- revenue: np.ndarray,
- avg_assets: np.ndarray,
- avg_equity: np.ndarray,
- ) -> np.ndarray:
- """ROE = Net_Profit_Margin × Asset_Turnover × Equity_Multiplier (DuPont)."""
- margin = FinancialMetricEngine.compute_net_profit_margin(pat, revenue)
- turnover = FinancialMetricEngine.compute_asset_turnover(revenue, avg_assets)
- multiplier = FinancialMetricEngine.compute_equity_multiplier(avg_assets, avg_equity)
- return margin * turnover * multiplier
-
- @staticmethod
- def compute_free_cash_flow(ocf: np.ndarray, capex: np.ndarray) -> np.ndarray:
- return ocf - capex
-
-
-# ── TensorBuilder ────────────────────────────────────────────────────────────
-
-class TensorBuilder:
- """
- Builds the (T, N, D) tensors for CUTS+ from financial data and news.
-
- KEY ARCHITECTURE DECISIONS
- --------------------------
- 1. Discovery nodes only — STRATEGIC_NODES excluded (algebraic identities).
- Use FinancialMetricEngine post-inference.
-
- 2. No interpolation — sparse observations + delta_t encoding.
- CUTS+ was designed for irregular time-series; interpolation leaks
- future information backward.
-
- 3. Expanding-window normalization — no future leakage.
-
- 4. Change-space features (FIX #1) — tensor is (T, N, 7):
- [level, qoq, yoy, acceleration, volatility, surprise, delta_t]
- CUTS+ discovers causality on changes/shocks, not accounting levels.
-
- 5. Release-date alignment (FIX #2) — fundamentals.reported_date used
- for timeline placement, NOT the fiscal period-end date.
-
- 6. EMA event persistence (FIX #3) — decayed EMA replaces additive
- accumulation to prevent embedding norm explosion.
-
- 7. Lag tensors (FIX #4) — lag-1, lag-3, lag-6 windows concatenated
- → final shape (T, N, 28).
-
- 8. Node attribution weights (FIX #5) — articles carry affected_nodes
- list with (node_name, weight) pairs for probabilistic distribution.
-
- 9. Surprise modeling (FIX #6) — channel 5 = z-score of current QoQ
- versus rolling 8-quarter history.
-
- 10. Financial DAG priors (FIX #7) — adjacency_mask returned alongside
- tensors; pass to CUTS+ as a structural prior.
-
- 11. Regime tensor (FIX #8) — (T, R) conditioning variables returned.
-
- 12. Dual embeddings (FIX #9) — 1792-D per article (sentiment + semantic).
-
- 13. Explicit horizon clipping (FIX #10) — _MAX_EVENT_HORIZON hard caps.
-
- freq="monthly": T ≈ 48 months / 4 years.
- freq="quarterly": T ≈ 16 — underdetermined for N=23; not recommended.
- freq="daily": T ≈ 1044 — 99%+ zeros; not recommended.
- """
-
- # Tech tensor channels (FIX #1)
- _TECH_BASE_CHANNELS = 7 # [level, qoq, yoy, accel, vol, surprise, delta_t]
- _LAG_WINDOWS = [0, 1, 3, 6] # 0 = current; 1/3/6 = lag in months
- _TECH_TOTAL_CHANNELS = _TECH_BASE_CHANNELS * len(_LAG_WINDOWS) # 28
-
- def __init__(self, symbols: Optional[List[str]] = None):
- self.nodes = DISCOVERY_NODES
- self.node_to_idx = {node: i for i, node in enumerate(self.nodes)}
- logger.info(
- f"TensorBuilder initialised with {len(self.nodes)} discovery nodes "
- f"({len(DERIVED_NODES)} derived nodes excluded from discovery)."
- )
-
- # ── Adjacency mask builder ─────────────────────────────────────────────
-
- def build_adjacency_mask(self) -> np.ndarray:
- """
- FIX #7 — Returns a (N, N) binary float32 mask where mask[i, j] = 1
- means 'node i is allowed to causally influence node j'.
- Pass to CUTS+ as prior_mask / structural_prior.
- """
- N = len(self.nodes)
- mask = np.zeros((N, N), dtype=np.float32)
- for src, targets in ALLOWED_CAUSAL_DIRECTIONS.items():
- if src not in self.node_to_idx:
- continue
- src_idx = self.node_to_idx[src]
- for tgt in targets:
- if tgt not in self.node_to_idx:
- continue
- mask[src_idx, self.node_to_idx[tgt]] = 1.0
- logger.info(
- f"Adjacency mask built: {int(mask.sum())} allowed edges "
- f"out of {N * N} possible."
- )
- return mask
-
- # ── Main build method ──────────────────────────────────────────────────
-
- def build(
- self,
- fundamentals: pd.DataFrame,
- causal_news: List[Dict[str, Any]],
- causal_policy: List[Dict[str, Any]],
- causal_gov: List[Dict[str, Any]],
- text_embeddings: Dict[str, np.ndarray], # article_url → sparse event vector
- start: str,
- end: str,
- freq: str = "monthly",
- macro_signals: Optional[pd.DataFrame] = None, # FIX #8 — optional regime inputs
- ) -> Tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray, np.ndarray, np.ndarray]:
- """
- Returns
- -------
- data_tech (T, N, 28) — 7 change-space channels × 4 lag windows.
- mask_tech (T, N, 1) — 1 at actual report dates, 0 otherwise.
- data_text (T, N, 84) — sparse event vectors × 3 modalities, EMA decay.
- mask_text (T, N, 3) — modality credibility weights.
- data_regime (T, R) — macro regime conditioning variables.
- adjacency_mask (N, N) — financial DAG structural prior for CUTS+.
-
- Fundamentals dataframe must have:
- - columns for each node (e.g. "Revenue", "COGS", ...)
- - a "reported_date" column (FIX #2) — the actual earnings release date.
- If absent, falls back to index (fiscal period end) with a warning.
- """
- # ── Strip timezone from index ──────────────────────────────────────
- if fundamentals.index.tz is not None:
- fundamentals.index = fundamentals.index.tz_localize(None)
-
- # ── Master timeline ────────────────────────────────────────────────
- if freq == "monthly":
- timeline = pd.date_range(start=start, end=end, freq="MS")
- logger.info(f"Using MONTHLY timeline: {len(timeline)} periods.")
- elif freq == "quarterly":
- timeline = pd.date_range(start=start, end=end, freq="QS")
- logger.info(f"Using QUARTERLY timeline: {len(timeline)} periods.")
- else:
- timeline = pd.bdate_range(start=start, end=end)
- logger.info(f"Using DAILY trading-day timeline: {len(timeline)} periods.")
-
- T = len(timeline)
- N = len(self.nodes)
- R = len(REGIME_FEATURES)
- F = self._TECH_BASE_CHANNELS
-
- logger.info(f"T={T}, N={N}, base_features={F}, total_tech_channels={F * len(self._LAG_WINDOWS)}")
- if T < N:
- logger.warning(
- f"T={T} < N={N}. Causal graph is underdetermined. "
- f"Switch to freq='monthly' (T~48) for reliable discovery."
- )
-
- # ── FIX #2 — Resolve reported_date column ──────────────────────────
- if "reported_date" in fundamentals.columns:
- reported_dates = pd.to_datetime(fundamentals["reported_date"]).dt.tz_localize(None)
- logger.info("Using reported_date column for timeline alignment (FIX #2).")
- else:
- reported_dates = pd.Series(fundamentals.index, index=fundamentals.index)
- logger.warning(
- "reported_date column not found in fundamentals. Falling back to fiscal "
- "period-end dates. This may introduce up to ~8 weeks of forward leakage per "
- "quarter. Add a reported_date column to fix this."
- )
-
- # ── Allocate tensors ───────────────────────────────────────────────
- # Base tech tensor before lag concatenation
- data_tech_base = np.zeros((T, N, F), dtype=np.float32)
- mask_tech = np.zeros((T, N, 1), dtype=np.float32)
- data_text = np.zeros((T, N, _TEXT_DIM), dtype=np.float32)
- mask_text = np.zeros((T, N, 3), dtype=np.float32)
- data_regime = np.zeros((T, R), dtype=np.float32)
-
- # ── Fill technical data ────────────────────────────────────────────
- logger.info("Populating technical data (change-space features, reported-date aligned)...")
-
- for node_idx, node in enumerate(self.nodes):
- if node not in fundamentals.columns:
- continue
-
- # Align raw values to reported dates (FIX #2)
- node_raw_vals = fundamentals[node].dropna()
- node_rep_dates = reported_dates.loc[node_raw_vals.index]
-
- if node_raw_vals.empty:
- continue
-
- # ── Mark true report dates (using reported_date) ──────────────
- for rep_date in node_rep_dates:
- t_idx = timeline.get_indexer([rep_date], method='nearest')[0]
- if 0 <= t_idx < T:
- mask_tech[t_idx, node_idx, 0] = 1.0
-
- # ── Build quarterly series aligned to fiscal period end for
- # computing growth rates (growth rates use fiscal ordering,
- # but signals are *placed* at reported_date).
- node_quarterly = pd.Series(
- node_raw_vals.values, index=node_raw_vals.index
- ).sort_index()
-
- # FIX #1 — Compute multi-scale dynamics on quarterly series
- qoq = node_quarterly.pct_change(periods=1)
- yoy = node_quarterly.pct_change(periods=4)
- accel = qoq.diff()
- vol = qoq.rolling(4, min_periods=2).std()
-
- # FIX #6 — Surprise: z-score of current QoQ vs rolling 8Q history
- roll_mean = qoq.rolling(8, min_periods=2).mean()
- roll_std = qoq.rolling(8, min_periods=2).std().replace(0.0, 1e-6)
- surprise = (qoq - roll_mean) / roll_std
-
- # Expanding-window normalisation for level channel (no future leakage)
- node_aligned = pd.Series(index=timeline, dtype=float)
- for fiscal_date, rep_date, val in zip(
- node_raw_vals.index, node_rep_dates, node_raw_vals.values
- ):
- t_idx = timeline.get_indexer([rep_date], method='nearest')[0]
- if 0 <= t_idx < T:
- node_aligned.iloc[t_idx] = val
-
- exp_mean = node_aligned.expanding().mean().ffill().bfill().fillna(0.0)
- exp_std = node_aligned.expanding().std().ffill().bfill().fillna(1.0).replace(0.0, 1.0)
- level_norm = ((node_aligned - exp_mean) / exp_std).ffill().bfill().fillna(0.0)
-
- # ── Map quarterly dynamics onto the monthly timeline ───────────
- # For each quarterly observation, place its dynamic values at the
- # reported_date timestep, then forward-fill (causal carry).
- def _map_to_timeline(quarterly_series: pd.Series, fill_val: float = 0.0) -> np.ndarray:
- out = pd.Series(index=timeline, dtype=float)
- for fiscal_date, rep_date, val in zip(
- node_raw_vals.index, node_rep_dates, quarterly_series.reindex(node_raw_vals.index).values
- ):
- t_idx = timeline.get_indexer([rep_date], method='nearest')[0]
- if 0 <= t_idx < T and not np.isnan(val):
- out.iloc[t_idx] = val
- return out.ffill().bfill().fillna(fill_val).values.astype(np.float32)
-
- data_tech_base[:, node_idx, 0] = level_norm.values.astype(np.float32)
- data_tech_base[:, node_idx, 1] = _map_to_timeline(qoq)
- data_tech_base[:, node_idx, 2] = _map_to_timeline(yoy)
- data_tech_base[:, node_idx, 3] = _map_to_timeline(accel)
- data_tech_base[:, node_idx, 4] = _map_to_timeline(vol)
- data_tech_base[:, node_idx, 5] = _map_to_timeline(surprise)
-
- # ── Channel 6: delta_t (days since last reported observation / 365) ──
- last_report_date = None
- for t_step, ts in enumerate(timeline):
- if mask_tech[t_step, node_idx, 0] == 1.0:
- last_report_date = ts
- if last_report_date is not None:
- data_tech_base[t_step, node_idx, 6] = (ts - last_report_date).days / 365.0
- else:
- data_tech_base[t_step, node_idx, 6] = 3.0 # sentinel: >3 years unseen
-
- # ── FIX #4 — Lag tensor concatenation ─────────────────────────────
- # Concatenate lag-0 (current), lag-1, lag-3, lag-6 along feature axis.
- # Roll-over boundary windows are zeroed out to avoid circular artefacts.
- lag_arrays = [data_tech_base]
- for lag in [1, 3, 6]:
- lagged = np.roll(data_tech_base, lag, axis=0)
- lagged[:lag] = 0.0
- lag_arrays.append(lagged)
-
- data_tech = np.concatenate(lag_arrays, axis=-1) # (T, N, 28)
- logger.info(
- f"Technical tensor built: {data_tech.shape} "
- f"({F} base channels × {len(self._LAG_WINDOWS)} lag windows)"
- )
-
- # Variance diagnostics (level channel)
- node_variances = np.var(data_tech[:, :, 0], axis=0)
- low_var_nodes = [self.nodes[i] for i in range(N) if node_variances[i] < 0.01]
- if low_var_nodes:
- logger.warning(f"Low-variance nodes (level < 0.01 std²): {low_var_nodes}")
- logger.info(
- f"Level channel variance — mean={node_variances.mean():.4f}, "
- f"min={node_variances.min():.4f}, max={node_variances.max():.4f}"
- )
-
- # ── Fill text data ─────────────────────────────────────────────────
- # FIX #3 — EMA decay (replaces additive accumulation).
- # FIX #5 — Probabilistic node attribution via affected_nodes list.
- # FIX #10 — Hard horizon caps per modality.
- logger.info("Populating multi-modal text data (EMA decay, node attribution weights)...")
-
- def populate_modality_with_ema_decay(
- causal_events: List[Dict[str, Any]],
- offset_dim: int,
- mask_idx: int,
- decay_lambda: float,
- max_horizon: int,
- ) -> int:
- """
- Propagates article embeddings forward in time using a decayed EMA.
-
- FIX #3: alpha = exp(-λ·dt); new = alpha*old + (1-alpha)*emb
- This keeps embedding norms bounded regardless of article count.
-
- FIX #5: Each article may carry an 'affected_nodes' list of
- (node_name, weight) tuples. The embedding contribution is
- scaled by node_weight before the EMA update. Falls back to
- the legacy 'affected_node' single-string key.
-
- FIX #10: Propagation stops at min(1e-4 threshold, max_horizon months).
- """
- event_count = 0
- for article in causal_events:
- # FIX #5 — resolve node attribution
- affected_nodes: List[Tuple[str, float]]
- if "affected_nodes" in article:
- affected_nodes = [
- (n, float(w)) for n, w in article["affected_nodes"]
- if n in self.node_to_idx
- ]
- elif article.get("affected_node") in self.node_to_idx:
- affected_nodes = [(article["affected_node"], 1.0)]
- else:
- continue
-
- url = article.get("url") or article.get("link")
- emb = text_embeddings.get(url)
- if emb is None:
- continue
-
- pub_date = pd.to_datetime(
- article.get("published", datetime.now())
- ).tz_localize(None)
- t_pub = timeline.get_indexer([pub_date], method="nearest")[0]
- if not (0 <= t_pub < T):
- continue
-
- article_weight = float(
- article.get("weighted_score", article.get("credibility_weight", 1.0))
- )
-
- for node_name, node_weight in affected_nodes:
- node_idx = self.node_to_idx[node_name]
- weighted_emb = node_weight * emb # scale by attribution weight
-
- for future_t in range(t_pub, T):
- dt = future_t - t_pub
-
- # FIX #10 — hard horizon cap
- if dt > max_horizon:
- break
-
- # FIX #3 — EMA decay
- alpha = float(np.exp(-decay_lambda * dt))
- if alpha < 1e-4:
- break # negligible — stop early
-
- start = offset_dim
- end = offset_dim + _COMBINED_DIM
- existing = data_text[future_t, node_idx, start:end]
- data_text[future_t, node_idx, start:end] = (
- alpha * existing + (1.0 - alpha) * weighted_emb
- )
-
- # Track max credibility weight across overlapping events
- mask_text[future_t, node_idx, mask_idx] = max(
- float(mask_text[future_t, node_idx, mask_idx]),
- article_weight * node_weight * alpha,
- )
-
- event_count += 1
- return event_count
-
- news_count = populate_modality_with_ema_decay(
- causal_news,
- offset_dim=0,
- mask_idx=0,
- decay_lambda=_DECAY_LAMBDA["news"],
- max_horizon=_MAX_EVENT_HORIZON["news"],
- )
- policy_count = populate_modality_with_ema_decay(
- causal_policy,
- offset_dim=_COMBINED_DIM,
- mask_idx=1,
- decay_lambda=_DECAY_LAMBDA["policy"],
- max_horizon=_MAX_EVENT_HORIZON["policy"],
- )
- gov_count = populate_modality_with_ema_decay(
- causal_gov,
- offset_dim=_COMBINED_DIM * 2,
- mask_idx=2,
- decay_lambda=_DECAY_LAMBDA["gov"],
- max_horizon=_MAX_EVENT_HORIZON["gov"],
- )
- text_count = news_count + policy_count + gov_count
- text_density = float(np.mean(mask_text > 0))
- logger.info(
- f"Text tensor built: {news_count} news, {policy_count} policy, {gov_count} gov events. "
- f"Mask density: {text_density:.4f}"
- )
-
- # ── FIX #8 — Regime tensor ─────────────────────────────────────────
- logger.info("Populating regime tensor...")
- if macro_signals is not None:
- # Caller may supply a DataFrame with columns matching REGIME_FEATURES,
- # indexed by date. We reindex to our timeline and forward-fill.
- if macro_signals.index.tz is not None:
- macro_signals = macro_signals.copy()
- macro_signals.index = macro_signals.index.tz_localize(None)
- for r_idx, feat in enumerate(REGIME_FEATURES):
- if feat in macro_signals.columns:
- aligned = macro_signals[feat].reindex(timeline, method="ffill").fillna(0.0)
- data_regime[:, r_idx] = aligned.values.astype(np.float32)
- else:
- logger.warning(
- "No macro_signals DataFrame provided. Regime tensor will be all zeros. "
- "Pass macro_signals= with columns for: " + ", ".join(REGIME_FEATURES)
- )
-
- # ── FIX #7 — Build adjacency mask ─────────────────────────────────
- adjacency_mask = self.build_adjacency_mask()
-
- # ── Save debug artefacts ───────────────────────────────────────────
- np.save(os.path.join(DEBUG_DIR, "data_tech.npy"), data_tech)
- np.save(os.path.join(DEBUG_DIR, "mask_tech.npy"), mask_tech)
- np.save(os.path.join(DEBUG_DIR, "data_text.npy"), data_text)
- np.save(os.path.join(DEBUG_DIR, "mask_text.npy"), mask_text)
- np.save(os.path.join(DEBUG_DIR, "data_regime.npy"), data_regime)
- np.save(os.path.join(DEBUG_DIR, "adjacency_mask.npy"), adjacency_mask)
- with open(os.path.join(DEBUG_DIR, "tensor_meta.json"), "w") as f:
- json.dump({
- "freq": freq,
- "T": T,
- "N": N,
- "tech_channels": self._TECH_TOTAL_CHANNELS,
- "text_dim": _TEXT_DIM,
- "regime_features": REGIME_FEATURES,
- "text_count": text_count,
- }, f, indent=2)
- logger.info(f"Saved all tensors and metadata to {DEBUG_DIR}")
-
- return data_tech, mask_tech, data_text, mask_text, data_regime, adjacency_mask
-
-
-# ── Standalone test ──────────────────────────────────────────────────────────
-
-def main():
- """
- Self-contained test for EmbeddingService and TensorBuilder.
- Exercises all 10 fixes with dummy data.
- """
- logging.basicConfig(level=logging.INFO)
- log = logging.getLogger("tensor_builder_test")
-
- start_str = "2020-01-01"
- end_str = "2023-12-31"
- dates = pd.date_range(start=start_str, end=end_str, freq="Q")
-
- # FIX #2 — fundamentals now carry a reported_date column
- # (typically 4–8 weeks after fiscal quarter end)
- fundamentals = pd.DataFrame(index=dates)
- for node in ["Revenue", "COGS", "PAT", "Inventory", "Total_Debt"]:
- fundamentals[node] = np.random.randn(len(dates)).cumsum() + 100
- # Simulate ~45-day reporting lag
- fundamentals["reported_date"] = dates + pd.DateOffset(days=45)
-
- # FIX #5 — causal events now carry affected_nodes with weights
- causal_news = [
- {
- "affected_nodes": [
- ("Revenue", 0.9),
- ("Inventory", 0.4),
- ],
- "url": "http://example.com/news1",
- "published": "2021-05-15",
- "title": "Company revenue skyrockets",
- "description": "The company saw a massive increase in revenue this quarter.",
- "weighted_score": 0.8,
- }
- ]
- causal_policy = [
- {
- "affected_nodes": [
- ("Inventory", 0.8),
- ("COGS", 0.6),
- ("Total_Debt", 0.3),
- ],
- "url": "http://example.com/policy1",
- "published": "2022-02-10",
- "title": "New import tax policy announced",
- "content": "Government increases import taxes affecting supply chains.",
- "credibility_weight": 0.9,
- }
- ]
- causal_gov = []
-
- # FIX #8 — dummy macro regime signals
- timeline = pd.date_range(start=start_str, end=end_str, freq="MS")
- macro_signals = pd.DataFrame(index=timeline)
- macro_signals["bull_market"] = (np.random.randn(len(timeline)).cumsum() > 0).astype(float)
- macro_signals["bear_market"] = 1.0 - macro_signals["bull_market"]
- macro_signals["high_inflation"] = (np.random.rand(len(timeline)) > 0.7).astype(float)
- macro_signals["tightening_cycle"] = (np.random.rand(len(timeline)) > 0.6).astype(float)
- macro_signals["earnings_expansion"] = (np.random.rand(len(timeline)) > 0.5).astype(float)
- macro_signals["earnings_contraction"] = 1.0 - macro_signals["earnings_expansion"]
- macro_signals["high_volatility"] = (np.random.rand(len(timeline)) > 0.75).astype(float)
-
- # ── Test EmbeddingService ──────────────────────────────────────────────
- log.info("Testing EmbeddingService (dual model)...")
- emb_service = EmbeddingService()
- articles_to_embed = causal_news + causal_policy
- embeddings_by_ticker = emb_service.embed_news({"DUMMY": articles_to_embed})
-
- url_to_embedding: Dict[str, np.ndarray] = {}
- for art in embeddings_by_ticker.get("DUMMY", []):
- url = art.get("url") or art.get("link")
- if url and "embedding" in art:
- url_to_embedding[url] = np.array(art["embedding"])
- log.info(f"Generated {len(url_to_embedding)} dual embeddings ({_COMBINED_DIM}-D each).")
-
- # ── Test TensorBuilder ─────────────────────────────────────────────────
- log.info("Testing TensorBuilder...")
- builder = TensorBuilder()
- (
- data_tech,
- mask_tech,
- data_text,
- mask_text,
- data_regime,
- adjacency_mask,
- ) = builder.build(
- fundamentals=fundamentals,
- causal_news=causal_news,
- causal_policy=causal_policy,
- causal_gov=causal_gov,
- text_embeddings=url_to_embedding,
- start=start_str,
- end=end_str,
- freq="monthly",
- macro_signals=macro_signals,
- )
-
- log.info("--- Tensor Build Results ---")
- log.info(f"data_tech shape : {data_tech.shape} (expected T x N x 28)")
- log.info(f"mask_tech shape : {mask_tech.shape} (expected T x N x 1)")
- log.info(f"data_text shape : {data_text.shape} (expected T x N x {_TEXT_DIM})")
- log.info(f"mask_text shape : {mask_text.shape} (expected T x N x 3)")
- log.info(f"data_regime shape : {data_regime.shape} (expected T x {len(REGIME_FEATURES)})")
- log.info(f"adjacency_mask : {adjacency_mask.shape} (expected N x N)")
-
- # ── Assertions ─────────────────────────────────────────────────────────
- assert data_tech.shape[2] == 28, \
- f"data_tech must have 28 channels (7 features × 4 lags), got {data_tech.shape[2]}"
- assert data_text.shape[2] == _TEXT_DIM, \
- f"data_text must be {_TEXT_DIM}-D (1792 × 3 modalities), got {data_text.shape[2]}"
- assert data_regime.shape[1] == len(REGIME_FEATURES), \
- f"data_regime must have {len(REGIME_FEATURES)} regime features"
- assert adjacency_mask.shape == (len(builder.nodes), len(builder.nodes)), \
- "adjacency_mask shape mismatch"
-
- # Verify change-space channels are non-trivial
- qoq_variance = np.nanvar(data_tech[:, :, 1])
- assert qoq_variance > 0, "QoQ channel (ch 1) should have non-zero variance"
-
- # Verify EMA decay doesn't blow up (bounded norms — FIX #3)
- rev_idx = builder.node_to_idx.get("Revenue", -1)
- if rev_idx >= 0:
- news_emb_slice = data_text[:, rev_idx, :_COMBINED_DIM]
- max_norm = float(np.max(np.linalg.norm(news_emb_slice, axis=-1)))
- log.info(f"Max embedding norm at Revenue node (news channel): {max_norm:.4f}")
- assert max_norm < 1e4, "Embedding norm suspiciously large — check EMA update"
-
- nonzero_steps = int(np.sum(np.any(news_emb_slice != 0, axis=-1)))
- log.info(f"Non-zero news embedding timesteps for Revenue: {nonzero_steps}")
- assert nonzero_steps > 1, "EMA decay should propagate beyond publication timestep"
-
- log.info("All assertions passed. Test completed successfully.")
-
-
-if __name__ == "__main__":
- main()
\ No newline at end of file
diff --git a/singular_ticker_causal/test_single_ticker_causal_flow.py b/singular_ticker_causal/test_single_ticker_causal_flow.py
deleted file mode 100644
index ece9eecc3b9b5e9cbf904aec7ae7e94a190fb5f7..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/test_single_ticker_causal_flow.py
+++ /dev/null
@@ -1,999 +0,0 @@
-"""
-test_fundamental_causal_flow.py
-
-End-to-end integration test for the singular-ticker fundamental causal pipeline.
-
-Pipeline:
- 1. Load fundamentals (from debug_data cache or yfinance fetch)
- 2. Fetch news (quarterly GDELT + RSS via FundamentalNewsClient, cached)
- 3. Embed news (FinBERT via EmbeddingService, cached)
- 4. Build quarterly (T, N, D) tensors (FundamentalTensorBuilder)
- 5. Build DuPont structural prior (build_dupont_prior)
- 6. Run CUTS+ MultiCAD with prior constraint + sparsity penalty
- 7. Validate: sparse graph, prior-consistent edges discovered
-"""
-
-import os
-import sys
-import logging
-import json
-import argparse
-import numpy as np
-import pandas as pd
-import torch
-from typing import List
-from datetime import datetime, timedelta
-from copy import deepcopy
-from omegaconf import OmegaConf
-
-# Add backend to sys.path if running as a standalone script
-sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
-
-from singular_ticker_causal.data_sources import (
- Fetcher, NewsClient, GDELTClient, SEBIREG30Client
-)
-from singular_ticker_causal.utils.llm_client import LLMClient
-from singular_ticker_causal.services.schema import (
- INCOME_STATEMENT_NODES,
- BALANCE_SHEET_NODES,
- STRATEGIC_NODES,
-)
-IND_AS_NODES = INCOME_STATEMENT_NODES + BALANCE_SHEET_NODES + STRATEGIC_NODES
-from singular_ticker_causal.services import (
- BOCDDetector, EmbeddingService,
- MSSAQuantEngine, TensorBuilder,
-)
-from singular_ticker_causal.services.tensor_builder import DISCOVERY_NODES
-from singular_ticker_causal.algorithms.CUTS_PLUS.cuts_plus import main as cuts_plus_main
-from singular_ticker_causal.algorithms.CUTS_PLUS.lagged_graph import discover_lagged_graphs
-from singular_ticker_causal.algorithms.CUTS_PLUS.utils.logger import MyLogger
-from causal_hierarchy import LLMNewsDenoiser
-from causal_hierarchy.hhkd import decompose_bidirectional_flux
-from causal_hierarchy.network import build_hierarchical_network
-from causal_hierarchy.metadata import ExternalTickerMetadataProvider
-from shared.news import FCMGraphInferencer
-
-
-logging.basicConfig(
- level=logging.INFO,
- format="%(asctime)s [%(levelname)s] %(name)s: %(message)s"
-)
-logger = logging.getLogger("test_fundamental_causal_flow")
-
-# ────────────────────────────────────────────────────────────────────────────
-# CUTS+ hyperparameters — tuned for sparse, prior-guided causal discovery
-# on monthly fundamental data (N=36 nodes, T~48 months for 4 years, T/N~1.33).
-# ────────────────────────────────────────────────────────────────────────────
-_TENSOR_FREQ = "monthly" # Change here to invalidate tensor cache automatically
-
-_TEST_CFG = OmegaConf.create({
- "n_nodes": len(DISCOVERY_NODES),
- "input_step": 3, # 3 months lookback (one quarter)
- "batch_size": 8,
- "data_dim": 28, # 7 change-space channels x 4 lag windows
- "total_epoch": 20, # More epochs: sparsity needs time to bite
- "n_groups": len(DISCOVERY_NODES), # No grouping — N=23 primitive nodes
- "group_policy": "None",
- "supervision_policy": "masked_before_8",
- "fill_policy": "None",
- "show_graph_every": 5,
- "projector_output_dim": 16,
- # Strong DuPont prior — known accounting edges are seeded at +2 logits
- "lambda_d": 1.0,
- "data_pred": {
- "model": "multi_lstm",
- "pred_step": 1,
- "mlp_hid": 32,
- "gru_layers": 1,
- "shared_weights_decoder": False,
- "concat_h": True,
- "lr_data_start": 5e-3,
- "lr_data_end": 5e-4,
- "weight_decay": 1e-5,
- "prob": False,
- },
- "graph_discov": {
- # 2.0 is 2x stronger than the causal/ pipeline's 1.0 —
- # necessary because N=36 and T=48 is a tighter ratio than the
- # 50-ticker causal pipeline where T >> N.
- "lambda_s_start": 2.0,
- "lambda_s_end": 0.5,
- "lr_graph_start": 5e-4, # Slower graph lr: let data predictor warm up first
- "lr_graph_end": 5e-5,
- # Fast tau annealing: by epoch 10, tau < 0.1 → near-binary Gumbel samples
- "start_tau": 1.0,
- "end_tau": 0.02,
- },
- "causal_thres": "value_0.5",
-})
-
-DEBUG_DATA_DIR = os.path.join(os.path.dirname(__file__), "debug_data")
-CACHE_DIR = DEBUG_DATA_DIR # same directory — one source of truth
-
-
-def _adj_to_graph(adj_matrix, symbols: List[str], threshold: float = 0.5):
- """
- Convert the (N, N) CUTS+ adjacency matrix into nodes/links dicts
- compatible with the frontend Voronoi/network visualisation.
- """
- nodes = [{"id": sym, "label": sym} for sym in symbols]
- links = []
- n = len(symbols)
- for i in range(n):
- for j in range(n):
- if i != j and float(adj_matrix[i, j]) >= threshold:
- links.append({
- "source": symbols[i],
- "target": symbols[j],
- "score": round(float(adj_matrix[i, j]), 4),
- })
- return nodes, links
-
-
-def _save_inference_results(
- ticker: str,
- graph_density: float,
- assert_result: dict | None = None,
- inter_result: dict | None = None,
- cf_result: dict | None = None,
-) -> str:
- """
- Persist causal-query results to ``debug_data/_inference_results.json``
- so the Streamlit Voronoi dashboard can load them without re-running inference.
- Returns the path written.
- """
- payload: dict = {"graph_density": graph_density}
- if assert_result:
- payload["assert_identifiable"] = assert_result.get("identifiable")
- payload["assert_strategy"] = assert_result.get("strategy")
- payload["assert_ate"] = float(assert_result["ate"]) if "ate" in assert_result else None
- payload["assert_ci_95"] = str(assert_result.get("ci_95", "—"))
- if inter_result:
- ape = inter_result.get("ate_per_target", {})
- payload["inter_delta_interest_expense"] = float(ape["Interest_Expense"]) if "Interest_Expense" in ape else None
- payload["inter_delta_ebt"] = float(ape["EBT"]) if "EBT" in ape else None
- if cf_result:
- payload["cf_factual_outcome"] = float(cf_result["factual_outcome"]) if "factual_outcome" in cf_result else None
- payload["cf_counterfactual_outcome"] = float(cf_result["counterfactual_outcome"]) if "counterfactual_outcome" in cf_result else None
- payload["cf_ite"] = float(cf_result["ite"]) if "ite" in cf_result else None
- out_path = os.path.join(DEBUG_DATA_DIR, f"{ticker}_inference_results.json")
- with open(out_path, "w") as fh:
- json.dump(payload, fh, indent=2)
- return out_path
-
-
-def _launch_voronoi_dashboard(ticker: str) -> None:
- """
- Spawn the Streamlit Voronoi dashboard as a background subprocess.
- The dashboard reads cached .npy / .json files from debug_data/ and does
- NOT need the full Python inference stack at runtime.
- """
- import subprocess
- dashboard_path = os.path.join(os.path.dirname(__file__), "voronoi_dashboard.py")
- cmd = [
- sys.executable, "-m", "streamlit", "run", dashboard_path,
- "--server.headless=false",
- "--", "--ticker", ticker,
- ]
- logger.info("Launching Voronoi dashboard: %s", " ".join(cmd))
- subprocess.Popen(cmd, start_new_session=True)
- logger.info(
- "Voronoi dashboard started — open http://localhost:8501 in your browser."
- )
-
-
-def _build_timeline(start: str, end: str, freq: str, expected_len: int) -> pd.DatetimeIndex:
- if freq == "monthly":
- timeline = pd.date_range(start=start, end=end, freq="MS")
- elif freq == "quarterly":
- timeline = pd.date_range(start=start, end=end, freq="QS")
- else:
- timeline = pd.bdate_range(start=start, end=end)
- if len(timeline) == expected_len:
- return timeline
- return pd.date_range(start=start, end=end, periods=expected_len)
-
-
-def _load_denoised_cache_entries(ticker: str) -> list[dict]:
- import base64
-
- cache_path = os.path.join(DEBUG_DATA_DIR, f"{ticker}_denoised_news_cache.json")
- if not os.path.exists(cache_path):
- return []
- with open(cache_path) as fh:
- cache = json.load(fh)
- entries = []
- for key, payload in cache.items():
- if "embedding" not in payload or float(payload.get("score", 0.0)) < 0.4:
- continue
- _, _, timestamp = key.partition(":")
- entries.append(
- {
- "timestamp": timestamp,
- "nodes": payload.get("nodes") or [],
- "embedding": np.frombuffer(base64.b64decode(payload["embedding"]), dtype=np.float32),
- }
- )
- return entries
-
-
-def _build_denoised_tensor(
- ticker: str,
- start: str,
- end: str,
- freq: str,
- T: int,
- N: int,
-) -> tuple[np.ndarray | None, np.ndarray | None]:
- entries = _load_denoised_cache_entries(ticker)
- if not entries:
- return None, None
-
- emb_dim = int(entries[0]["embedding"].shape[0])
- data_denoised = np.zeros((T, N, emb_dim), dtype=np.float32)
- mask_denoised = np.zeros((T, N, 1), dtype=np.float32)
- timeline = _build_timeline(start, end, freq, T)
- node_to_idx = {node: idx for idx, node in enumerate(DISCOVERY_NODES)}
-
- for entry in entries:
- pub_date = pd.to_datetime(entry["timestamp"], errors="coerce")
- if pd.isna(pub_date):
- continue
- if getattr(pub_date, "tzinfo", None) is not None:
- pub_date = pub_date.tz_localize(None)
- t_idx = timeline.get_indexer([pub_date], method="nearest")[0]
- if not (0 <= t_idx < T):
- continue
- node_indices = [node_to_idx[node] for node in entry["nodes"] if node in node_to_idx]
- if not node_indices:
- node_indices = list(range(N))
- for node_idx in node_indices:
- data_denoised[t_idx, node_idx] = entry["embedding"]
- mask_denoised[t_idx, node_idx, 0] = 1.0
-
- if not np.any(mask_denoised):
- return None, None
- np.save(os.path.join(DEBUG_DATA_DIR, "data_denoised_news.npy"), data_denoised)
- np.save(os.path.join(DEBUG_DATA_DIR, "mask_denoised_news.npy"), mask_denoised)
- return data_denoised, mask_denoised
-
-
-def _infer_fcm_sequence(
- data_text: np.ndarray,
- data_tech: np.ndarray,
- num_lags: int = 3,
-) -> np.ndarray:
- L = min(num_lags, data_text.shape[0], data_tech.shape[0])
- if L < 1:
- raise ValueError("Cannot infer FCM sequence with no timesteps")
- C = torch.from_numpy(data_text[-L:]).float()
- P = torch.from_numpy(data_tech[-L:]).float()
- inferencer = FCMGraphInferencer(D_nodes=data_tech.shape[1], L=L, hidden_dim=64)
- inferencer.eval()
- with torch.no_grad():
- return torch.sigmoid(inferencer(C, P)).cpu().numpy().astype(np.float32)
-
-
-def test_end_to_end_flow(inference_only: bool = False):
- ticker = "RELIANCE"
- end_dt_obj = datetime.now()
- start_dt_obj = end_dt_obj - timedelta(days=365 * 4) # 4 years
-
- start_str = start_dt_obj.strftime("%Y-%m-%d")
- end_str = end_dt_obj.strftime("%Y-%m-%d")
-
- logger.info(f"--- Starting E2E Test for {ticker} ({start_str} → {end_str}) ---")
- os.makedirs(DEBUG_DATA_DIR, exist_ok=True)
- idx = {name: i for i, name in enumerate(DISCOVERY_NODES)}
- tensor_files = {
- "data_tech": os.path.join(DEBUG_DATA_DIR, "data_tech.npy"),
- "mask_tech": os.path.join(DEBUG_DATA_DIR, "mask_tech.npy"),
- "data_text": os.path.join(DEBUG_DATA_DIR, "data_text.npy"),
- "mask_text": os.path.join(DEBUG_DATA_DIR, "mask_text.npy"),
- "data_regime": os.path.join(DEBUG_DATA_DIR, "data_regime.npy"),
- "adjacency_mask": os.path.join(DEBUG_DATA_DIR, "adjacency_mask.npy"),
- }
- adj_save_path = os.path.join(DEBUG_DATA_DIR, f"{ticker}_adj_matrix.npy")
-
- if inference_only:
- logger.info("Inference-only mode: loading cached tensors and adjacency from debug_data/")
- required = [
- tensor_files["data_tech"],
- tensor_files["mask_tech"],
- tensor_files["data_text"],
- tensor_files["adjacency_mask"],
- adj_save_path,
- ]
- missing = [p for p in required if not os.path.exists(p)]
- if missing:
- raise FileNotFoundError(
- "Inference-only mode requires cached artifacts. Missing:\n"
- + "\n".join(missing)
- )
- data_tech = np.load(tensor_files["data_tech"])
- mask_tech = np.load(tensor_files["mask_tech"])
- data_text = np.load(tensor_files["data_text"])
- adjacency_mask = np.load(tensor_files["adjacency_mask"])
- adj_matrix = np.load(adj_save_path)
- logger.info(
- "Loaded inference artifacts: data_tech=%s, mask_tech=%s, data_text=%s, adjacency_mask=%s, adj=%s",
- data_tech.shape, mask_tech.shape, data_text.shape, adjacency_mask.shape, adj_matrix.shape
- )
- # ── Step 10: Causal Inference Queries (inference-only path) ─────────
- from singular_ticker_causal.causal_inference.causal_model import StructuralCausalModel
- from singular_ticker_causal.causal_inference.query_engine import CausalQueryEngine
-
- logger.info("10. Building SCM and running causal queries...")
- scm = StructuralCausalModel(
- nodes=DISCOVERY_NODES,
- adj=adj_matrix,
- adjacency_mask=adjacency_mask,
- prohibition_mask=adjacency_mask,
- data_tech=data_tech,
- mask_tech=mask_tech,
- threshold=0.5,
- lag=1,
- ).fit()
-
- graph_density = scm.density()
- logger.info(f" SCM graph density: {graph_density:.3f}")
- if graph_density > 0.4:
- logger.warning(
- " SCM density is high (>0.4), so effect estimates may be unstable."
- )
-
- engine = CausalQueryEngine(scm, data_tech=data_tech, data_text=data_text)
-
- assert_result: dict | None = None
- inter_result: dict | None = None
- cf_result: dict | None = None
-
- try:
- assert_result = engine.assert_edge("Revenue", "EBIT")
- logger.info(
- " Assertion Revenue->EBIT: identifiable=%s strategy=%s ate=%.6f ci_95=%s",
- assert_result["identifiable"],
- assert_result["strategy"],
- float(assert_result["ate"]),
- assert_result["ci_95"],
- )
- assert "ate" in assert_result and "ci_95" in assert_result and "strategy" in assert_result
- except Exception as exc:
- logger.warning(" Assertion Revenue->EBIT skipped: %s", exc)
-
- debt_idx = idx["Total_Debt"]
- debt_value = float(data_tech[-1, debt_idx, 0]) * 1.10
- try:
- inter_result = engine.intervene(
- "Total_Debt",
- value=debt_value,
- targets=["Interest_Expense", "EBT"],
- horizon=5,
- )
- logger.info(
- " Intervention Total_Debt +10%%: ΔInterest_Expense(t+1)=%.6f, ΔEBT(t+1)=%.6f",
- inter_result["ate_per_target"]["Interest_Expense"],
- inter_result["ate_per_target"]["EBT"],
- )
- assert set(inter_result["predicted_values"].keys()) == {"Interest_Expense", "EBT"}
- assert len(inter_result["predicted_values"]["Interest_Expense"]) == 5
- except Exception as exc:
- logger.warning(" Intervention query skipped: %s", exc)
-
- revenue_idx = idx["Revenue"]
- try:
- cf_result = engine.counterfactual(
- observed_t=-1,
- treatment="Revenue",
- cf_value=float(data_tech[-1, revenue_idx, 0]) * 1.05,
- target="PAT",
- )
- logger.info(
- " Counterfactual PAT: factual=%.6f cf=%.6f ite=%.6f",
- cf_result["factual_outcome"],
- cf_result["counterfactual_outcome"],
- cf_result["ite"],
- )
- assert "factual_outcome" in cf_result and "counterfactual_outcome" in cf_result and "ite" in cf_result
- except Exception as exc:
- logger.warning(" Counterfactual query skipped: %s", exc)
-
- # ── Persist results & optionally launch Voronoi dashboard ────────────
- results_path = _save_inference_results(
- ticker=ticker,
- graph_density=graph_density,
- assert_result=assert_result,
- inter_result=inter_result,
- cf_result=cf_result,
- )
- logger.info(" Inference results persisted to %s", results_path)
-
- logger.info("--- Inference-only run completed successfully ---")
- return
-
- else:
- # ── Step 1: Fundamentals (cache-first) ───────────────────────────────────
- fundamentals_file = os.path.join(DEBUG_DATA_DIR, f"{ticker}.NS_processed_fundamentals.csv")
- if os.path.exists(fundamentals_file):
- logger.info(f"1. [CACHE] Loading fundamentals from {fundamentals_file}")
- fundamentals = pd.read_csv(fundamentals_file, index_col=0, parse_dates=True)
- else:
- logger.info("1. [FETCH] Fetching fundamentals via yfinance...")
- fetcher = Fetcher(ticker)
- fundamentals = fetcher.fetch(start_str, end_str)
- # Explicitly save to cache after fetch
- fundamentals.to_csv(fundamentals_file)
- logger.info(f" Saved fundamentals to {fundamentals_file}")
-
- # MSSA Denoising
- logger.info("1a. [RUN] Applying MSSA denoising (fundamentals)...")
- engine = MSSAQuantEngine(window_size=4, variance_threshold=0.90)
- if len(fundamentals) >= engine.window_size * 2:
- numeric_cols = fundamentals.select_dtypes(include="number").columns.tolist()
- sub = fundamentals[numeric_cols].ffill().bfill().fillna(0)
- engine.fit(sub)
- denoised = engine.extract_signal()
- fundamentals = denoised.reindex(fundamentals.index).combine_first(fundamentals)
- logger.info(f" MSSA denoising complete. Denoised shape: {denoised.shape}")
- else:
- logger.warning(f" Skipping MSSA — not enough rows ({len(fundamentals)} < {engine.window_size * 2})")
-
- # BOCD Changepoint Detection
- logger.info("1b. [RUN] Running Change Point Detection (PELT)...")
- detector = BOCDDetector(model="rbf", min_size=4, penalty_scale=2.0, use_log_returns=False)
- changepoints = detector.detect_dataframe_changepoints(fundamentals)
- num_cp = sum(len(v) for v in changepoints.values())
- logger.info(f" Detected {num_cp} total structural breaks across all nodes.")
-
- logger.info(f" Shape: {fundamentals.shape} | Columns: {list(fundamentals.columns[:6])}...")
- assert not fundamentals.empty, "Fundamentals DataFrame is empty"
- assert "Revenue" in fundamentals.columns, "Revenue node missing from fundamentals"
-
- # ── Step 2: Raw news corpus (cache-first) ────────────────────────────────
- raw_news_file = os.path.join(CACHE_DIR, f"{ticker}_raw_news_corpus.json")
- if os.path.exists(raw_news_file):
- logger.info(f"2. [CACHE] Loading raw news from {raw_news_file}")
- with open(raw_news_file, "r") as f:
- raw_articles = json.load(f)
- logger.info(f" Loaded {len(raw_articles)} articles.")
- else:
- logger.info("2. [FETCH] Fetching news corpus (quarterly GDELT windows)...")
- news_client = NewsClient(ticker)
- raw_articles = news_client.fetch(start_str, end_str, include_pulse=True, include_te=True)
- # Explicitly save to cache after fetch
- with open(raw_news_file, "w") as f:
- json.dump(raw_articles, f, indent=2)
- logger.info(f" Saved raw news to {raw_news_file}")
- logger.info(f" Fetched {len(raw_articles)} articles.")
-
- assert raw_articles, "No news articles available — stopping execution."
-
- # ── Step 3: Denoised news encoding (cache-first) ─────────────────────────
- causal_results_file = os.path.join(DEBUG_DATA_DIR, f"{ticker}_causal_results.json")
- if os.path.exists(causal_results_file):
- logger.info(f"3. [CACHE] Loading causal results from {causal_results_file}")
- with open(causal_results_file, "r") as f:
- causal_articles = json.load(f)
- else:
- logger.info("3. [RUN] Encoding and filtering causal news events using LLM denoiser...")
- try:
- _llm_client = LLMClient()
- _emb_service_for_denoiser = EmbeddingService()
-
- def _embed_one(text: str) -> np.ndarray:
- """Embed a single article text using EmbeddingService."""
- results = _emb_service_for_denoiser.embed_news({ticker: [{"text": text, "url": "_inline_"}]})
- articles_out = results.get(ticker, [])
- if articles_out and "embedding" in articles_out[0]:
- return np.array(articles_out[0]["embedding"], dtype=np.float32)
- raise ValueError("EmbeddingService returned no embedding for inline text")
-
- denoiser = LLMNewsDenoiser(
- llm_client=_llm_client,
- embedding_fn=_embed_one,
- threshold=0.4,
- cache_dir=DEBUG_DATA_DIR,
- )
- # filter_and_embed returns {timestamp: weighted_embedding}; we also
- # keep the full raw_articles list for downstream steps that need
- # the article dicts (e.g. policy/governance steps).
- denoised_embeddings = denoiser.filter_and_embed(raw_articles, ticker)
- # Build causal_articles as those raw articles whose timestamp passed
- # the denoiser threshold.
- passed_timestamps = set(denoised_embeddings.keys())
- causal_articles = [
- a for a in raw_articles
- if (a.get("published") or a.get("date") or a.get("publishedAt")) in passed_timestamps
- ]
- if not causal_articles:
- logger.warning(
- " LLM denoiser filtered ALL articles for %s. "
- "Falling back to raw articles to avoid empty pipeline.",
- ticker,
- )
- causal_articles = raw_articles
- except Exception as _denoiser_exc:
- logger.warning(
- " LLM denoiser unavailable (%s). "
- "Falling back to raw articles (no filtering).",
- _denoiser_exc,
- )
- causal_articles = raw_articles
-
- with open(causal_results_file, "w") as f:
- json.dump(causal_articles, f, indent=2)
-
- logger.info(
- f" Causal news events: {len(causal_articles)} across nodes: "
- f"{sorted({a.get('affected_node') for a in causal_articles} - {None})}"
- )
- assert causal_articles, f"No causal news events extracted for {ticker} — stopping."
-
- # ── Step 3b: Policy encoding (cache-first) ─────────────────────────────────
- policy_results_file = os.path.join(DEBUG_DATA_DIR, f"{ticker}_policy_results.json")
- if os.path.exists(policy_results_file):
- logger.info(f"3b. [CACHE] Loading policy results from {policy_results_file}")
- with open(policy_results_file, "r") as f:
- causal_policy = json.load(f)
- else:
- logger.info("3b. [RUN] Fetching and encoding policy events (GDELT + RSS)...")
- policy_client = GDELTClient()
- llm_client = LLMClient()
-
- # We use a more specific query to avoid GDELT's broad query rate limits
- policy_query = "RBI OR SEBI OR 'Union Budget India'"
-
- all_raw_policy = []
- chunk_start = start_dt_obj
- while chunk_start < end_dt_obj:
- chunk_end = min(chunk_start + timedelta(days=89), end_dt_obj) # Larger chunks, fewer calls
- logger.info(f" Policy chunk: {chunk_start.date()} → {chunk_end.date()}")
-
- # Since GDELT DOC API only has history for the last 90-120 days from NOW,
- # older chunks will return empty or error. We still call it for completeness,
- # but we use 89 day chunks to respect the client constraint.
- raw_chunk = policy_client.fetch(policy_query, chunk_start, chunk_end)
- all_raw_policy.extend(raw_chunk)
- chunk_start = chunk_end
- import time as _time
- _time.sleep(2) # Longer sleep to respect GDELT
-
- # Supplementary: Pull policy from RSS feeds as well (high reliability)
- logger.info(" [RUN] Searching RSS feeds for supplementary policy coverage...")
- from singular_ticker_causal.data_sources.news_client import search_rss_feeds, livemint_feeds, cnbc18_feeds, other_feeds, _filter_by_window
- rss_policy = (
- search_rss_feeds(livemint_feeds, "RBI OR SEBI", "LiveMint-Policy") +
- search_rss_feeds(cnbc18_feeds, "RBI OR SEBI", "CNBC18-Policy") +
- search_rss_feeds(other_feeds, "RBI OR SEBI", "Other-Policy")
- )
- rss_policy_filtered = _filter_by_window(rss_policy, start_dt_obj, end_dt_obj)
- all_raw_policy.extend(rss_policy_filtered)
-
- causal_policy = all_raw_policy
- with open(policy_results_file, "w") as f:
- json.dump(causal_policy, f, indent=2)
-
- logger.info(f" Causal policy events: {len(causal_policy)}")
- assert causal_policy, "No causal policy events extracted — stopping."
-
- # ── Step 3c: Governance encoding (cache-first) ─────────────────────────────
- gov_results_file = os.path.join(DEBUG_DATA_DIR, f"{ticker}_gov_results.json")
- if os.path.exists(gov_results_file):
- logger.info(f"3c. [CACHE] Loading governance results from {gov_results_file}")
- with open(gov_results_file, "r") as f:
- causal_gov = json.load(f)
- else:
- logger.info("3c. [RUN] Fetching and encoding governance events...")
- gov_client = SEBIREG30Client()
- raw_gov = gov_client.fetch(ticker, start_str, end_str)
- llm_client = LLMClient()
- causal_gov = raw_gov
- with open(gov_results_file, "w") as f:
- json.dump(causal_gov, f, indent=2)
-
- logger.info(f" Causal governance events: {len(causal_gov)}")
- assert causal_gov, "No causal governance events extracted — stopping."
-
- # ── Step 4: FinBERT embeddings (cache-first, incremental) ─────────────────
- emb_cache_file = os.path.join(CACHE_DIR, f"{ticker}_embedding_cache.json")
- url_to_embedding: dict = {}
-
- if os.path.exists(emb_cache_file):
- logger.info(f"4. [CACHE] Loading embeddings from {emb_cache_file}")
- with open(emb_cache_file, "r") as f:
- raw_cache = json.load(f)
- url_to_embedding = {url: np.array(vec) for url, vec in raw_cache.items()}
- logger.info(f" Loaded {len(url_to_embedding)} cached embeddings.")
- else:
- logger.info("4. [RUN] Generating FinBERT embeddings (no cache found)...")
-
- # Only embed articles whose URL is not yet in the cache (incremental updates)
- urls_needed = {
- a.get("url") or a.get("link")
- for a in causal_articles + causal_policy + causal_gov
- if (a.get("url") or a.get("link")) and
- (a.get("url") or a.get("link")) not in url_to_embedding
- }
-
- if urls_needed:
- logger.info(f" [RUN] Embedding {len(urls_needed)} new articles not in cache...")
- articles_to_embed = [
- a for a in causal_articles + causal_policy + causal_gov
- if (a.get("url") or a.get("link")) in urls_needed
- ]
- emb_service = EmbeddingService()
- embeddings_by_ticker = emb_service.embed_news({ticker: articles_to_embed})
- for art_with_emb in embeddings_by_ticker.get(ticker, []):
- url = art_with_emb.get("url") or art_with_emb.get("link")
- if "embedding" in art_with_emb and url:
- url_to_embedding[url] = np.array(art_with_emb["embedding"])
- # Persist updated cache
- with open(emb_cache_file, "w") as f:
- json.dump({u: v.tolist() for u, v in url_to_embedding.items()}, f)
- logger.info(f" Cache updated: {len(url_to_embedding)} total embeddings saved.")
- else:
- logger.info(f" All {len(url_to_embedding)} embeddings already in cache — skipping FinBERT.")
-
- assert len(url_to_embedding) > 0 or not causal_articles, \
- "No embeddings and causal articles exist — embedding pipeline broken."
-
- # ── Step 5: Tensor Builder (cache-first, freq-aware invalidation) ─────────
- tensor_files = {
- "data_tech": os.path.join(DEBUG_DATA_DIR, "data_tech.npy"),
- "mask_tech": os.path.join(DEBUG_DATA_DIR, "mask_tech.npy"),
- "data_text": os.path.join(DEBUG_DATA_DIR, "data_text.npy"),
- "mask_text": os.path.join(DEBUG_DATA_DIR, "mask_text.npy"),
- "data_regime": os.path.join(DEBUG_DATA_DIR, "data_regime.npy"),
- "adjacency_mask": os.path.join(DEBUG_DATA_DIR, "adjacency_mask.npy"),
- }
- tensor_meta_file = os.path.join(DEBUG_DATA_DIR, "tensor_meta.json")
- all_tensors_cached = all(os.path.exists(p) for p in tensor_files.values())
-
- # Invalidate cache if freq has changed since last build
- if all_tensors_cached and os.path.exists(tensor_meta_file):
- import json as _json
- with open(tensor_meta_file) as f:
- cached_meta = _json.load(f)
- if cached_meta.get("freq") != _TENSOR_FREQ:
- logger.info(
- f"5. [STALE] Tensor cache built with freq='{cached_meta.get('freq')}' "
- f"but current freq='{_TENSOR_FREQ}'. Rebuilding..."
- )
- all_tensors_cached = False
-
- if all_tensors_cached:
- logger.info("5. [CACHE] Loading tensors from saved .npy files in debug_data/")
- data_tech = np.load(tensor_files["data_tech"])
- mask_tech = np.load(tensor_files["mask_tech"])
- data_text = np.load(tensor_files["data_text"])
- mask_text = np.load(tensor_files["mask_text"])
- data_regime = np.load(tensor_files["data_regime"])
- adjacency_mask = np.load(tensor_files["adjacency_mask"])
- else:
- logger.info(f"5. [RUN] Building {_TENSOR_FREQ} T × N × D tensors with linear interpolation...")
- builder = TensorBuilder()
- data_tech, mask_tech, data_text, mask_text, data_regime, adjacency_mask = builder.build(
- fundamentals, causal_articles, causal_policy, causal_gov, url_to_embedding, start_str, end_str,
- freq=_TENSOR_FREQ,
- )
-
- T, N, D_tech = data_tech.shape
- D_text = data_text.shape[2]
- tech_density = float(np.mean(mask_tech > 0))
- text_density = float(np.mean(mask_text > 0))
-
- logger.info(f" data_tech: {data_tech.shape} (channels: value / obs_mask / delta_t)")
- logger.info(f" data_text: {data_text.shape} mask density: {text_density:.4f}")
- logger.info(f" obs density: {tech_density:.4f}")
-
- # N is now DISCOVERY_NODES count (primitive only — STRATEGIC_NODES excluded)
- assert N == len(DISCOVERY_NODES), f"Expected {len(DISCOVERY_NODES)} discovery nodes, got {N}"
- assert D_tech == 28, f"Expected D_tech=28 (7 channels × 4 lag windows), got {D_tech}"
- assert D_text in (4608, 5376), f"Expected D_text=4608 or 5376, got {D_text}"
- assert tech_density > 0.001, \
- f"Technical mask density too low ({tech_density:.4f}). " \
- f"Check that fundamentals index aligns with the monthly timeline."
-
- data_denoised, mask_denoised = _build_denoised_tensor(
- ticker=ticker,
- start=start_str,
- end=end_str,
- freq=_TENSOR_FREQ,
- T=T,
- N=N,
- )
- if data_denoised is not None:
- logger.info(
- " data_denoised_news: %s mask density: %.4f",
- data_denoised.shape,
- float(np.mean(mask_denoised > 0)),
- )
- else:
- logger.warning(" No denoised-news tensor available; CUTS+ will use raw text only.")
-
- R = data_regime.shape[1]
- assert data_regime.shape == (T, R), f"data_regime shape mismatch: {data_regime.shape}"
- assert adjacency_mask.shape == (N, N), f"adjacency_mask shape mismatch: {adjacency_mask.shape}"
- logger.info(f" data_regime: {data_regime.shape} (T x {R} regime features)")
- logger.info(f" adjacency_mask: {adjacency_mask.shape} ({int(adjacency_mask.sum())} allowed edges)")
-
- # ── Step 6: DuPont Structural Prior ─────────────────────────────────────
- # NOTE: The prior is now built over DISCOVERY_NODES only (primitive nodes).
- # Edges involving derived KPI nodes (ROE, Gross_Profit, etc.) are excluded
- # because those nodes are no longer part of the discovery graph.
- logger.info("6. Building DuPont structural prior (discovery nodes only)...")
- G_prior = adjacency_mask # use the richer FIX #7 prior directly
- prior_edge_count = int(G_prior.sum())
- logger.info(f" Prior shape: {G_prior.shape} | Known edges: {prior_edge_count}")
-
- # idx maps discovery node names to their positions
- idx = {name: i for i, name in enumerate(DISCOVERY_NODES)}
- # Verify a known primitive-to-primitive prior edge (Revenue → COGS direction via cost structure)
- # Only assert edges that exist between DISCOVERY_NODES
- assert G_prior.shape == (len(DISCOVERY_NODES), len(DISCOVERY_NODES)), \
- f"Prior shape mismatch: expected ({len(DISCOVERY_NODES)}, {len(DISCOVERY_NODES)}), got {G_prior.shape}"
-
- # ── Step 6b: FCMGraphInferencer lag sequence ──────────────────────────────
- # Phase 4: produce G_{1:L} before CUTS+ so the sequence can be cached and
- # inspected as the temporal counterpart to the static discovered graph.
- fcm_cache_path = os.path.join(DEBUG_DATA_DIR, f"{ticker}_fcm_graph_sequence.npy")
- if os.path.exists(fcm_cache_path):
- fcm_bundle_graphs = np.load(fcm_cache_path)
- logger.info("6b. [CACHE] FCM graph sequence loaded: %s", fcm_bundle_graphs.shape)
- else:
- logger.info("6b. [RUN] Inferring FCM graph sequence with GRU graph inferencer...")
- fcm_source = data_denoised if data_denoised is not None else data_text
- fcm_bundle_graphs = _infer_fcm_sequence(fcm_source, data_tech, num_lags=3)
- np.save(fcm_cache_path, fcm_bundle_graphs)
- logger.info(" FCM graph sequence saved: %s", fcm_bundle_graphs.shape)
-
- # ── Step 7: CUTS+ MultiCAD ───────────────────────────────────────────────
- logger.info("7. Running CUTS+ MultiCAD (quarterly, prior-guided, sparse)...")
- cfg = deepcopy(_TEST_CFG)
- cfg.n_nodes = N
-
- log_dir = os.path.join(os.path.dirname(__file__), "test_logs")
- os.makedirs(log_dir, exist_ok=True)
- cuts_log = MyLogger(log_dir=log_dir, stdout=True, tensorboard=False)
-
- adj_matrix = cuts_plus_main(
- data=data_tech,
- mask=mask_tech,
- true_cm=None,
- opt=cfg,
- log=cuts_log,
- device="cpu",
- text_data=data_text,
- text_mask=mask_text,
- G_prior=G_prior,
- denoised_news=data_denoised,
- denoised_mask=mask_denoised,
- )
-
- logger.info(f" Adjacency matrix: {adj_matrix.shape}")
- assert adj_matrix.shape == (N, N)
-
- # Save the adjacency matrix for inference reuse
- adj_save_path = os.path.join(DEBUG_DATA_DIR, f"{ticker}_adj_matrix.npy")
- np.save(adj_save_path, adj_matrix)
- logger.info(f" Saved adjacency matrix to {adj_save_path}")
-
- # ── Step 7b: Lag-Dependent FCM Graph Sequence ────────────────────────────
- # (Phase 4 / Phase 9): Discover G_{1:L} by chaining CUTS+ across lags.
- if os.path.exists(fcm_cache_path):
- logger.info("7b. [CACHE] Loading FCM lagged graph sequence from cache...")
- fcm_bundle_graphs = np.load(fcm_cache_path)
- logger.info(f" FCM graph sequence shape: {fcm_bundle_graphs.shape}")
- else:
- logger.info("7b. [RUN] Discovering lag-dependent FCM graph sequence (L=3 lags)...")
- _fcm_cfg = deepcopy(cfg)
- fcm_bundle = discover_lagged_graphs(
- data=data_tech,
- mask=mask_tech,
- opt=_fcm_cfg,
- log=cuts_log,
- device="cpu",
- text_data=data_text,
- text_mask=mask_text,
- G_prior=G_prior,
- num_lags=3,
- reducer="max",
- )
- fcm_bundle_graphs = fcm_bundle.graphs # (L, N, N)
- np.save(fcm_cache_path, fcm_bundle_graphs)
- logger.info(
- f" FCM graph sequence: {fcm_bundle_graphs.shape} "
- f"| input_steps={fcm_bundle.input_steps.tolist()}"
- )
-
- # ── Step 7c: HHKD Node Potential Ranking ────────────────────────────────
- # (Phase 1 / Phase 9): Decompose the discovered adjacency into a scalar
- # potential phi that ranks nodes from upstream driver to downstream sink.
- phi_cache_path = os.path.join(DEBUG_DATA_DIR, f"{ticker}_node_potentials.npy")
- macro_graph_cache_path = os.path.join(DEBUG_DATA_DIR, f"{ticker}_sector_macro_graph.json")
- logger.info("7c. [RUN] HHKD decomposition — computing node potentials phi...")
- hhkd_result = decompose_bidirectional_flux(adj_matrix)
- np.save(phi_cache_path, hhkd_result.phi)
- logger.info(
- f" phi saved to {phi_cache_path} | "
- f"top-3 upstream nodes: {hhkd_result.upstream_ranking[:3]}"
- )
- logger.info(
- f" Gradient component norm: {np.linalg.norm(hhkd_result.J_gradient):.4f} | "
- f"Residual cyclic norm: {np.linalg.norm(hhkd_result.J_residual):.4f}"
- )
-
- # Build and cache the sector macro graph via HierarchicalNetwork
- try:
- _metadata_provider = ExternalTickerMetadataProvider(
- cache_dir=os.path.join(DEBUG_DATA_DIR, "metadata_cache"),
- static_overrides={ticker: {"sector": "unknown", "industry": "unknown"}},
- )
- _hier_network = build_hierarchical_network(
- symbols=DISCOVERY_NODES,
- ticker_adj=adj_matrix,
- ticker_metadata=_metadata_provider.resolve(DISCOVERY_NODES),
- )
- import json as _json_mod
- with open(macro_graph_cache_path, "w") as _mf:
- _json_mod.dump(
- {
- "sector_labels": _hier_network.sector_labels,
- "macro_graph": _hier_network.macro_graph.tolist(),
- "n_cross_edges": len(_hier_network.cross_level_edges),
- },
- _mf,
- indent=2,
- )
- logger.info(f" Sector macro graph saved to {macro_graph_cache_path}")
- except Exception as _hier_exc:
- logger.warning(" Could not build sector macro graph: %s", _hier_exc)
-
- # ── Step 8: Graph Quality Validation ────────────────────────────────────
- logger.info("8. Validating causal graph quality...")
-
- # ── Raw adjacency diagnostics (detect saturation before thresholding) ──
- adj_flat = adj_matrix.flatten()
- logger.info(
- f" Raw adjacency stats: "
- f"min={adj_flat.min():.4f}, max={adj_flat.max():.4f}, "
- f"mean={adj_flat.mean():.4f}, std={adj_flat.std():.4f}"
- )
- near_one = float(np.mean(adj_flat > 0.95))
- near_zero = float(np.mean(adj_flat < 0.05))
- logger.info(
- f" Saturation check: {near_one*100:.1f}% of edges > 0.95, "
- f"{near_zero*100:.1f}% of edges < 0.05"
- )
- if near_one > 0.8:
- logger.warning(
- " ⚠ SATURATION DETECTED: >80% of adj values are near 1.0. "
- "The model has not converged to a sparse solution. "
- "Try increasing lambda_s_start or total_epoch."
- )
-
- nodes, links = _adj_to_graph(adj_matrix, DISCOVERY_NODES, threshold=0.5)
- max_edges = N * (N - 1)
- density_pct = len(links) / max_edges * 100
-
- logger.info(
- f" threshold=0.5 → {len(links)} edges / {max_edges} possible "
- f"({density_pct:.1f}% dense)"
- )
-
- # Must NOT be fully connected
- assert len(links) < max_edges, (
- f"Fully-connected graph ({len(links)} edges = 100% dense). "
- f"Raw adj stats: min={adj_flat.min():.4f}, max={adj_flat.max():.4f}. "
- f"Increase lambda_s_start or total_epoch."
- )
-
- # Identify prior-consistent edges
- prior_recovered = [
- lnk for lnk in links
- if lnk["source"] in idx and lnk["target"] in idx
- and G_prior[idx[lnk["source"]], idx[lnk["target"]]] == 1
- ]
- logger.info(
- f" Prior-consistent edges recovered: "
- f"{len(prior_recovered)} / {prior_edge_count} known"
- )
-
- if links:
- top5 = sorted(links, key=lambda x: x["score"], reverse=True)[:5]
- logger.info(" Top-5 edges by score:")
- for lnk in top5:
- prior_mark = "✓" if G_prior[idx.get(lnk["source"], 0), idx.get(lnk["target"], 0)] else " "
- logger.info(
- f" [{prior_mark}] {lnk['source']} → {lnk['target']} "
- f"score={lnk['score']:.4f}"
- )
-
- # ── Step 10: Causal Inference Queries ───────────────────────────────────
- from singular_ticker_causal.causal_inference.causal_model import StructuralCausalModel
- from singular_ticker_causal.causal_inference.query_engine import CausalQueryEngine
-
- logger.info("10. Building SCM and running causal queries...")
- scm = StructuralCausalModel(
- nodes=DISCOVERY_NODES,
- adj=adj_matrix,
- adjacency_mask=adjacency_mask,
- prohibition_mask=adjacency_mask,
- data_tech=data_tech,
- mask_tech=mask_tech,
- threshold=0.5,
- lag=1,
- ).fit()
-
- graph_density = scm.density()
- logger.info(f" SCM graph density: {graph_density:.3f}")
- if graph_density > 0.4:
- logger.warning(
- " SCM density is high (>0.4), so effect estimates may be unstable."
- )
-
- engine = CausalQueryEngine(scm, data_tech=data_tech, data_text=data_text)
-
- # Assertion: Revenue -> EBIT
- assert_result = engine.assert_edge("Revenue", "EBIT")
- logger.info(
- " Assertion Revenue->EBIT: identifiable=%s strategy=%s ate=%.6f ci_95=%s",
- assert_result["identifiable"],
- assert_result["strategy"],
- float(assert_result["ate"]),
- assert_result["ci_95"],
- )
- assert "ate" in assert_result and "ci_95" in assert_result and "strategy" in assert_result
-
- # Intervention: raise Total_Debt by 10%
- debt_idx = idx["Total_Debt"]
- debt_value = float(data_tech[-1, debt_idx, 0]) * 1.10
- inter_result = engine.intervene(
- "Total_Debt",
- value=debt_value,
- targets=["Interest_Expense", "EBT"],
- horizon=5,
- )
- logger.info(
- " Intervention Total_Debt +10%%: ΔInterest_Expense(t+1)=%.6f, ΔEBT(t+1)=%.6f",
- inter_result["ate_per_target"]["Interest_Expense"],
- inter_result["ate_per_target"]["EBT"],
- )
- assert set(inter_result["predicted_values"].keys()) == {"Interest_Expense", "EBT"}
- assert len(inter_result["predicted_values"]["Interest_Expense"]) == 5
-
- # Counterfactual: if Revenue had been 5% higher at t=-1, what about PAT?
- revenue_idx = idx["Revenue"]
- cf_result = engine.counterfactual(
- observed_t=-1,
- treatment="Revenue",
- cf_value=float(data_tech[-1, revenue_idx, 0]) * 1.05,
- target="PAT",
- )
- logger.info(
- " Counterfactual PAT: factual=%.6f cf=%.6f ite=%.6f",
- cf_result["factual_outcome"],
- cf_result["counterfactual_outcome"],
- cf_result["ite"],
- )
- assert "factual_outcome" in cf_result and "counterfactual_outcome" in cf_result and "ite" in cf_result
-
- logger.info("--- E2E Test Completed Successfully ---")
-
-
-if __name__ == "__main__":
- parser = argparse.ArgumentParser(description="Run singular ticker causal E2E flow")
- parser.add_argument(
- "--inference-only",
- action="store_true",
- help="Skip Steps 1-9 and run only Step 10 by loading cached .npy artifacts from debug_data/",
- )
- parser.add_argument(
- "--voronoi",
- action="store_true",
- help="After inference, auto-launch the Streamlit Voronoi dashboard at http://localhost:8501.",
- )
- args = parser.parse_args()
- test_end_to_end_flow(inference_only=args.inference_only)
- if args.voronoi:
- _launch_voronoi_dashboard(ticker="RELIANCE")
diff --git a/singular_ticker_causal/utils/config.py b/singular_ticker_causal/utils/config.py
deleted file mode 100644
index dea949b88d2cbcfd87e0a888dc1f91dee634ce6d..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/utils/config.py
+++ /dev/null
@@ -1,34 +0,0 @@
-import os
-from dotenv import load_dotenv
-
-# Load environment variables from .env file
-load_dotenv()
-
-class Config:
- # Provider Settings
- LLM_PROVIDER = os.environ.get("LLM_PROVIDER", "nvidia") # "nvidia" | "fireworks"
-
- # Fireworks LLM Settings
- FIREWORKS_API_KEY = os.environ.get("FIREWORKS_API_KEY")
- FIREWORKS_BASE_URL = os.environ.get("FIREWORKS_BASE_URL", "https://api.fireworks.ai/inference/v1")
- # Primary LLM model (Fireworks)
- FIREWORKS_PRIMARY_MODEL = os.environ.get(
- "FIREWORKS_PRIMARY_MODEL",
- "accounts/fireworks/models/minimax-m2p5"
- )
- FIREWORKS_FALLBACK_MODEL = os.environ.get(
- "FIREWORKS_FALLBACK_MODEL",
- "accounts/fireworks/models/llama-v3p1-8b-instruct"
- )
-
- # NVIDIA LLM Settings
- NVIDIA_API_KEY = os.environ.get("NVIDIA_API_KEY")
- NVIDIA_BASE_URL = os.environ.get("NVIDIA_BASE_URL", "https://integrate.api.nvidia.com/v1")
- NVIDIA_PRIMARY_MODEL = os.environ.get(
- "NVIDIA_PRIMARY_MODEL",
- "nvidia/nemotron-3-nano-30b-a3b"
- )
- NVIDIA_FALLBACK_MODEL = os.environ.get(
- "NVIDIA_FALLBACK_MODEL",
- "meta/llama-3.1-8b-instruct"
- )
diff --git a/singular_ticker_causal/utils/llm_client.py b/singular_ticker_causal/utils/llm_client.py
deleted file mode 100644
index 99177400544cbe286bc9ee853484775be442cc02..0000000000000000000000000000000000000000
--- a/singular_ticker_causal/utils/llm_client.py
+++ /dev/null
@@ -1,205 +0,0 @@
-import json
-import logging
-import os
-import random
-import time
-import requests
-from requests.adapters import HTTPAdapter
-from urllib3.util.retry import Retry
-
-from .config import Config
-
-
-logger = logging.getLogger("causal-chain-api")
-
-
-def _build_session() -> requests.Session:
- """Build a requests Session with conservative urllib3 retries (quick network blips only)."""
- session = requests.Session()
- retries = Retry(
- total=2,
- backoff_factor=0.3,
- status_forcelist=[429, 500, 502, 503, 504],
- allowed_methods=["POST"],
- )
- session.mount("https://", HTTPAdapter(pool_connections=10, pool_maxsize=10, max_retries=retries))
- return session
-
-
-class LLMClient:
- """Wrapper around Fireworks/NVIDIA APIs for consistent interaction across services.
-
- Retry strategy (two-tier):
- 1. urllib3 makes 2 quick attempts per request for transient blips.
- 2. Application-level loop retries up to `max_app_retries` times with
- exponential back-off + jitter.
- 3. If all retries on the primary model fail, it transparently switches
- to the configured fallback model and repeats the loop once more.
- """
-
- FIREWORKS_BASE_URL = f"{Config.FIREWORKS_BASE_URL}/chat/completions"
- NVIDIA_BASE_URL = f"{Config.NVIDIA_BASE_URL}/chat/completions"
-
- def __init__(self, api_key: str = None, base_url: str = None, model: str = None, provider: str = None):
- self.provider = provider or Config.LLM_PROVIDER
-
- if self.provider == "nvidia":
- self.api_key = api_key or Config.NVIDIA_API_KEY or os.environ.get("NVIDIA_API_KEY")
- self.base_url = base_url or self.NVIDIA_BASE_URL
- self.model = model or Config.NVIDIA_PRIMARY_MODEL
- self.fallback_model = Config.NVIDIA_FALLBACK_MODEL
- else: # fireworks
- self.api_key = api_key or Config.FIREWORKS_API_KEY or os.environ.get("FIREWORKS_API_KEY")
- self.base_url = base_url or self.FIREWORKS_BASE_URL
- self.model = model or Config.FIREWORKS_PRIMARY_MODEL
- self.fallback_model = Config.FIREWORKS_FALLBACK_MODEL
-
- if not self.api_key:
- raise ValueError(f"API key for provider '{self.provider}' must be provided")
-
- self.session = _build_session()
-
- # ── internal helpers ──────────────────────────────────────────────────────
-
- def _headers(self) -> dict:
- return {
- "Content-Type": "application/json",
- "Authorization": f"Bearer {self.api_key}",
- }
-
- def _consume_stream(self, response: requests.Response) -> str:
- """Consumes an SSE stream from Fireworks/OpenAI API and returns full content."""
- full_content = ""
- for line in response.iter_lines():
- if not line:
- continue
- line_str = line.decode("utf-8")
- if line_str.startswith("data: "):
- data_str = line_str[6:].strip()
- if data_str == "[DONE]":
- break
- try:
- chunk = json.loads(data_str)
- if "choices" in chunk and len(chunk["choices"]) > 0:
- delta = chunk["choices"][0].get("delta", {})
- full_content += delta.get("content", "")
- except json.JSONDecodeError:
- continue
- return full_content
-
- def _post_with_retry(self, payload: dict, *, max_app_retries: int = 4) -> requests.Response:
- """
- POST payload to LLM API, retrying with exponential backoff + jitter.
- On complete failure, tries the fallback model once more before raising.
- """
- is_stream = payload.get("stream", False)
- headers = self._headers()
-
- models_to_try = [payload["model"]]
- if payload["model"] != self.fallback_model:
- models_to_try.append(self.fallback_model)
-
- last_exc: Exception | None = None
-
- for model_candidate in models_to_try:
- candidate_payload = {**payload, "model": model_candidate}
- if model_candidate != payload["model"]:
- logger.warning(
- f"[LLMClient] Primary model exhausted retries. "
- f"Switching to fallback model: {model_candidate}"
- )
-
- for attempt in range(max_app_retries):
- try:
- response = self.session.post(
- self.base_url,
- headers=headers,
- json=candidate_payload,
- timeout=120,
- stream=is_stream,
- )
- response.raise_for_status()
- return response
- except Exception as exc:
- last_exc = exc
- if attempt < max_app_retries - 1:
- # Exponential backoff with jitter: base * 2^attempt ± 20% noise
- base_wait = 2 * (2 ** attempt)
- jitter = random.uniform(-0.2 * base_wait, 0.2 * base_wait)
- wait_time = max(1.0, base_wait + jitter)
- logger.warning(
- f"[LLMClient] model={model_candidate} attempt {attempt + 1}/{max_app_retries} "
- f"failed: {exc}. Retrying in {wait_time:.1f}s..."
- )
- time.sleep(wait_time)
- else:
- logger.error(
- f"[LLMClient] model={model_candidate} failed after "
- f"{max_app_retries} attempts: {exc}"
- )
-
- raise last_exc
-
- # ── public API ────────────────────────────────────────────────────────────
-
- def chat(self, messages: list[dict], temperature: float = 0.6, max_tokens: int = 4096) -> str:
- """Standard chat completion returning text content."""
- payload = {
- "model": self.model,
- "max_tokens": max_tokens,
- "top_p": 1,
- "presence_penalty": 0,
- "frequency_penalty": 0,
- "temperature": temperature,
- "messages": messages,
- }
- if self.provider == "fireworks":
- payload["top_k"] = 40
- if max_tokens > 4096:
- payload["stream"] = True
-
- response = self._post_with_retry(payload)
-
- if payload.get("stream"):
- return self._consume_stream(response)
- return response.json()["choices"][0]["message"]["content"]
-
- def chat_json(self, messages: list[dict], temperature: float = 0.6, max_tokens: int = 4096) -> dict:
- """Chat completion enforcing JSON output format."""
- payload = {
- "model": self.model,
- "max_tokens": max_tokens,
- "top_p": 1,
- "presence_penalty": 0,
- "frequency_penalty": 0,
- "temperature": temperature,
- "messages": messages,
- "response_format": {"type": "json_object"},
- }
- if self.provider == "fireworks":
- payload["top_k"] = 40
- if max_tokens > 4096:
- payload["stream"] = True
-
- response = self._post_with_retry(payload)
-
- if payload.get("stream"):
- content = self._consume_stream(response)
- else:
- content = response.json()["choices"][0]["message"]["content"]
-
- try:
- return json.loads(content)
- except json.JSONDecodeError as e:
- # Fallback for models that sometimes wrap JSON in markdown blocks
- if "```" in content:
- import re
- match = re.search(r'```(?:json)?\s*(.*?)\s*```', content, re.DOTALL)
- if match:
- try:
- return json.loads(match.group(1))
- except json.JSONDecodeError:
- pass
-
- logger.error(f"[LLMClient] Failed to parse JSON response: {content[:200]}...")
- raise e