You cannot select more than 25 topics
Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
211 lines
9.2 KiB
Python
211 lines
9.2 KiB
Python
"""Build continuation-focused snapshots from existing event-day snapshots."""
|
|
from __future__ import annotations
|
|
|
|
import datetime as dt
|
|
import json
|
|
from decimal import Decimal
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
import pyarrow as pa
|
|
import pyarrow.parquet as pq
|
|
|
|
from libs.common.config import get_settings
|
|
from libs.common.time_utils import trading_days_between, utc_now
|
|
from libs.export.snapshot_export import _rows_to_table, _temporal_split
|
|
from libs.features.market_features import compute_market_features
|
|
from libs.labeler.label_generator import _compute_labels_from_bars, _pct_return
|
|
from libs.oracle_client.models import PriceBar
|
|
|
|
|
|
def _parse_date(raw: Any) -> dt.date | None:
|
|
if isinstance(raw, dt.date):
|
|
return raw
|
|
if isinstance(raw, str):
|
|
try:
|
|
return dt.date.fromisoformat(raw)
|
|
except ValueError:
|
|
return None
|
|
return None
|
|
|
|
|
|
def _iso_ts(date_value: dt.date) -> str:
|
|
return f"{date_value.isoformat()}T21:00:00+00:00"
|
|
|
|
|
|
def _to_price_bars(date_bars: dict[dt.date, dict[str, Any]]) -> list[PriceBar]:
|
|
rows: list[PriceBar] = []
|
|
for date_value in sorted(date_bars):
|
|
bar = date_bars[date_value]
|
|
rows.append(
|
|
PriceBar(
|
|
date=date_value.isoformat(),
|
|
open=float(bar["open"]),
|
|
high=float(bar["high"]),
|
|
low=float(bar["low"]),
|
|
close=float(bar["close"]),
|
|
volume=float(bar["volume"]),
|
|
)
|
|
)
|
|
return rows
|
|
|
|
|
|
def _continuation_entry_window(exec_date: dt.date, lookback_days: int) -> tuple[dt.date, dt.date] | None:
|
|
trading_days = trading_days_between(exec_date, exec_date + dt.timedelta(days=14))
|
|
signal_index = lookback_days
|
|
entry_index = lookback_days + 1
|
|
if len(trading_days) <= entry_index:
|
|
return None
|
|
return trading_days[signal_index], trading_days[entry_index]
|
|
|
|
|
|
def _build_continuation_rows_from_bars(
|
|
base_rows: list[dict[str, Any]],
|
|
bars_by_symbol: dict[str, dict[dt.date, dict[str, Any]]],
|
|
*,
|
|
lookback_days: int = 3,
|
|
) -> list[dict[str, Any]]:
|
|
transformed: list[dict[str, Any]] = []
|
|
price_bars_cache: dict[str, list[PriceBar]] = {}
|
|
|
|
for row in base_rows:
|
|
ticker = str(row.get("ticker") or row.get("symbol") or "").upper()
|
|
exec_date = _parse_date(row.get("entry_date") or row.get("execution_date"))
|
|
source_close_raw = row.get("event_close") or row.get("entry_price_est")
|
|
if not ticker or exec_date is None or source_close_raw in (None, 0):
|
|
continue
|
|
source_close = float(source_close_raw)
|
|
|
|
symbol_bars = bars_by_symbol.get(ticker)
|
|
if not symbol_bars or exec_date not in symbol_bars:
|
|
continue
|
|
|
|
dates = _continuation_entry_window(exec_date, lookback_days)
|
|
if dates is None:
|
|
continue
|
|
signal_date, entry_date = dates
|
|
if signal_date not in symbol_bars or entry_date not in symbol_bars:
|
|
continue
|
|
|
|
signal_bar = symbol_bars[signal_date]
|
|
signal_close = float(signal_bar["close"])
|
|
if source_close <= 0 or signal_close <= 0:
|
|
continue
|
|
drift_pct = (signal_close - source_close) / source_close
|
|
|
|
price_bars = price_bars_cache.get(ticker)
|
|
if price_bars is None:
|
|
price_bars = _to_price_bars(symbol_bars)
|
|
price_bars_cache[ticker] = price_bars
|
|
signal_features = compute_market_features(price_bars, signal_date.isoformat())
|
|
|
|
sorted_dates = sorted(symbol_bars)
|
|
entry_idx = sorted_dates.index(entry_date)
|
|
bars_from_entry = [symbol_bars[d] for d in sorted_dates[entry_idx:]]
|
|
if not bars_from_entry:
|
|
continue
|
|
entry_price = Decimal(str(bars_from_entry[0]["open"]))
|
|
forward_bars = bars_from_entry[1:]
|
|
|
|
label_status = "ok" if len(forward_bars) >= 20 else "truncated"
|
|
fwd_1d = _pct_return(entry_price, Decimal(str(bars_from_entry[1]["close"]))) if len(bars_from_entry) > 1 else None
|
|
lbl_3d = _compute_labels_from_bars(entry_price, forward_bars, 3)
|
|
lbl_5d = _compute_labels_from_bars(entry_price, forward_bars, 5)
|
|
lbl_10d = _compute_labels_from_bars(entry_price, forward_bars, 10)
|
|
lbl_20d = _compute_labels_from_bars(entry_price, forward_bars, 20)
|
|
|
|
new_row = dict(row)
|
|
original_event_id = str(row.get("event_id") or "")
|
|
new_row["event_id"] = f"{original_event_id}::cont_d{lookback_days}"
|
|
new_row["original_event_id"] = original_event_id
|
|
new_row["original_event_date"] = row.get("event_date")
|
|
new_row["original_reaction_date"] = row.get("reaction_date")
|
|
new_row["original_entry_date"] = row.get("entry_date")
|
|
new_row["original_event_close"] = source_close
|
|
new_row["event_date"] = signal_date.isoformat()
|
|
new_row["reaction_date"] = signal_date.isoformat()
|
|
new_row["entry_date"] = entry_date.isoformat()
|
|
new_row["entry_convention"] = "next_open_after_continuation_signal"
|
|
new_row["event_timestamp"] = _iso_ts(signal_date)
|
|
new_row["event_close"] = signal_close
|
|
new_row["entry_price"] = float(entry_price)
|
|
new_row["reaction_day_return"] = drift_pct
|
|
new_row["continuation_anchor_drift_pct"] = drift_pct
|
|
new_row["continuation_anchor_day_return"] = signal_features.get("reaction_day_return")
|
|
new_row["close_location"] = signal_features.get("close_location")
|
|
new_row["volume_ratio_20d"] = signal_features.get("volume_ratio_20d")
|
|
new_row["avg_dollar_volume_20d"] = signal_features.get("avg_dollar_volume_20d")
|
|
new_row["gap_size"] = signal_features.get("gap_size")
|
|
new_row["atr_14"] = signal_features.get("atr_14")
|
|
new_row["reaction_day_low"] = signal_features.get("reaction_day_low")
|
|
new_row["reaction_day_high"] = signal_features.get("reaction_day_high")
|
|
new_row["fwd_return_1d"] = float(fwd_1d) if fwd_1d is not None else None
|
|
new_row["fwd_return_3d"] = float(lbl_3d.get("fwd_return")) if lbl_3d.get("fwd_return") is not None else None
|
|
new_row["fwd_return_5d"] = float(lbl_5d.get("fwd_return")) if lbl_5d.get("fwd_return") is not None else None
|
|
new_row["fwd_return_10d"] = float(lbl_10d.get("fwd_return")) if lbl_10d.get("fwd_return") is not None else None
|
|
new_row["fwd_return_20d"] = float(lbl_20d.get("fwd_return")) if lbl_20d.get("fwd_return") is not None else None
|
|
new_row["mfe_3d"] = float(lbl_3d.get("mfe")) if lbl_3d.get("mfe") is not None else None
|
|
new_row["mae_3d"] = float(lbl_3d.get("mae")) if lbl_3d.get("mae") is not None else None
|
|
new_row["mfe_5d"] = float(lbl_5d.get("mfe")) if lbl_5d.get("mfe") is not None else None
|
|
new_row["mae_5d"] = float(lbl_5d.get("mae")) if lbl_5d.get("mae") is not None else None
|
|
new_row["mfe_10d"] = float(lbl_10d.get("mfe")) if lbl_10d.get("mfe") is not None else None
|
|
new_row["mae_10d"] = float(lbl_10d.get("mae")) if lbl_10d.get("mae") is not None else None
|
|
new_row["mfe_20d"] = float(lbl_20d.get("mfe")) if lbl_20d.get("mfe") is not None else None
|
|
new_row["mae_20d"] = float(lbl_20d.get("mae")) if lbl_20d.get("mae") is not None else None
|
|
new_row["label_status"] = label_status
|
|
transformed.append(new_row)
|
|
|
|
return transformed
|
|
|
|
|
|
async def export_continuation_snapshot_from_base(
|
|
*,
|
|
base_snapshot_dir: str | Path,
|
|
output_dir: str | Path,
|
|
snapshot_id: str,
|
|
lookback_days: int = 3,
|
|
) -> dict[str, Any]:
|
|
from libs.backtest.snapshot_store import SnapshotStore
|
|
|
|
base_path = Path(base_snapshot_dir)
|
|
rows: list[dict[str, Any]] = []
|
|
for split in ("train", "valid", "test"):
|
|
table = pq.read_table(base_path / f"{split}.parquet")
|
|
rows.extend(table.to_pylist())
|
|
|
|
symbols = sorted({str(r.get("ticker") or r.get("symbol") or "").upper() for r in rows if r.get("ticker") or r.get("symbol")})
|
|
dates = [_parse_date(r.get("entry_date")) for r in rows]
|
|
dates = [d for d in dates if d is not None]
|
|
if not symbols or not dates:
|
|
raise ValueError("Base snapshot has no symbols or dates")
|
|
|
|
settings = get_settings()
|
|
bars_by_symbol, _ = await SnapshotStore._fetch_price_data(
|
|
symbols,
|
|
(min(dates) - dt.timedelta(days=45), max(dates) + dt.timedelta(days=45)),
|
|
settings.stock_oracle_url,
|
|
)
|
|
transformed = _build_continuation_rows_from_bars(rows, bars_by_symbol, lookback_days=lookback_days)
|
|
splits = _temporal_split(transformed, "temporal_70_15_15")
|
|
|
|
out_path = Path(output_dir) / snapshot_id
|
|
out_path.mkdir(parents=True, exist_ok=True)
|
|
row_counts: dict[str, int] = {}
|
|
for split_name, split_rows in splits.items():
|
|
pq.write_table(_rows_to_table(split_rows), out_path / f"{split_name}.parquet")
|
|
row_counts[split_name] = len(split_rows)
|
|
|
|
manifest = {
|
|
"snapshot_id": snapshot_id,
|
|
"created_at_utc": utc_now().isoformat(),
|
|
"base_snapshot_id": base_path.name,
|
|
"transform": f"continuation_d{lookback_days}",
|
|
"split_policy": "temporal_70_15_15",
|
|
"row_counts": row_counts,
|
|
"total_rows": sum(row_counts.values()),
|
|
"output_dir": str(out_path),
|
|
}
|
|
(out_path / "manifest.json").write_text(json.dumps(manifest, indent=2))
|
|
return manifest
|
|
|