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.
262 lines
9.0 KiB
Python
262 lines
9.0 KiB
Python
"""Generate forward-return labels for parsed events using price bars."""
|
|
from __future__ import annotations
|
|
|
|
import datetime as dt
|
|
from decimal import Decimal
|
|
from typing import Any
|
|
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from libs.common.logging import get_logger
|
|
from libs.common.time_utils import filing_time_bucket as classify_time_bucket, next_trading_day, trading_days_between
|
|
from libs.labeler.reaction_date import compute_reaction_date
|
|
|
|
logger = get_logger(__name__)
|
|
|
|
LABEL_VERSION = "label-2.0.0"
|
|
|
|
_LOOK_AHEAD_DAYS = 25 # fetch enough trading days for 20D label computation
|
|
_R_FACTOR = 0.01 # 1R = 1% move (used for hit_pos/neg_1r labels)
|
|
|
|
|
|
def _safe_decimal(v: Any) -> Decimal | None:
|
|
try:
|
|
return Decimal(str(v)) if v is not None else None
|
|
except Exception:
|
|
return None
|
|
|
|
|
|
def _pct_return(entry: Decimal, exit_price: Decimal) -> Decimal:
|
|
if entry == 0:
|
|
return Decimal("0")
|
|
return (exit_price - entry) / entry
|
|
|
|
|
|
def _compute_labels_from_bars(
|
|
entry_price: Decimal,
|
|
bars: list[dict[str, Any]], # sorted ascending by date
|
|
n_days: int,
|
|
) -> dict[str, Any]:
|
|
"""Compute forward-return labels over `n_days` trading days of bars.
|
|
|
|
Args:
|
|
entry_price: Entry price (open of entry date).
|
|
bars: List of OHLCV dicts with keys: date, open, high, low, close.
|
|
n_days: Look-ahead horizon (3 or 5).
|
|
|
|
Returns:
|
|
Dict of computed label fields for the given horizon.
|
|
"""
|
|
window = bars[:n_days]
|
|
if not window:
|
|
return {}
|
|
|
|
closes = [_safe_decimal(b.get("close")) for b in window]
|
|
highs = [_safe_decimal(b.get("high")) for b in window]
|
|
lows = [_safe_decimal(b.get("low")) for b in window]
|
|
|
|
# Forward close return at day n
|
|
last_close = closes[-1]
|
|
fwd_return = _pct_return(entry_price, last_close) if last_close else None
|
|
|
|
# MFE (max favorable excursion): max high vs entry price
|
|
valid_highs = [h for h in highs if h is not None]
|
|
mfe = (max(valid_highs) - entry_price) / entry_price if valid_highs else None
|
|
|
|
# MAE (max adverse excursion): min low vs entry price
|
|
valid_lows = [lo for lo in lows if lo is not None]
|
|
mae = (min(valid_lows) - entry_price) / entry_price if valid_lows else None
|
|
|
|
# Hit +1R within n days
|
|
threshold_pos = entry_price * (1 + Decimal(str(_R_FACTOR)))
|
|
hit_pos = any(h is not None and h >= threshold_pos for h in highs)
|
|
|
|
# Hit -1R within n days
|
|
threshold_neg = entry_price * (1 - Decimal(str(_R_FACTOR)))
|
|
hit_neg = any(lo is not None and lo <= threshold_neg for lo in lows)
|
|
|
|
# Close up after n days
|
|
close_up = bool(last_close is not None and last_close > entry_price)
|
|
|
|
# Bars to MFE (index of max high)
|
|
bars_to_mfe: int | None = None
|
|
if valid_highs and mfe is not None:
|
|
max_high = max(valid_highs)
|
|
for i, h in enumerate(highs):
|
|
if h == max_high:
|
|
bars_to_mfe = i + 1
|
|
break
|
|
|
|
# Days to peak close
|
|
valid_close_idx = [(i, c) for i, c in enumerate(closes) if c is not None]
|
|
days_to_peak_close: int | None = None
|
|
if valid_close_idx:
|
|
peak_close_idx = max(valid_close_idx, key=lambda x: x[1])[0]
|
|
days_to_peak_close = peak_close_idx + 1
|
|
|
|
return {
|
|
"fwd_return": fwd_return,
|
|
"mfe": mfe,
|
|
"mae": mae,
|
|
"hit_pos_1r": hit_pos,
|
|
"hit_neg_1r": hit_neg,
|
|
"close_up": close_up,
|
|
"bars_to_mfe": bars_to_mfe,
|
|
"days_to_peak_close": days_to_peak_close,
|
|
}
|
|
|
|
|
|
async def generate_labels(
|
|
session: AsyncSession,
|
|
event: Any, # Event ORM model instance
|
|
price_svc: Any, # PriceService
|
|
ticker: str,
|
|
entry_convention: str = "next_open_after_reaction_close",
|
|
) -> Any:
|
|
"""Generate EventLabel for a single event.
|
|
|
|
Args:
|
|
session: Async DB session.
|
|
event: Event ORM instance (needs event_date, filed_at_utc, symbol_id).
|
|
price_svc: PriceService instance for fetching bars.
|
|
ticker: Trading ticker symbol.
|
|
entry_convention: How to determine entry price.
|
|
|
|
Returns:
|
|
EventLabel ORM instance (not yet added to session).
|
|
"""
|
|
from libs.db.models import EventLabel
|
|
|
|
# 1. Compute reaction_date
|
|
if event.filed_at_utc is not None:
|
|
filing_time_bucket = classify_time_bucket(event.filed_at_utc)
|
|
else:
|
|
filing_time_bucket = "unknown"
|
|
|
|
event_date: dt.date = event.event_date
|
|
reaction_date = compute_reaction_date(event_date, filing_time_bucket)
|
|
|
|
# 2. Compute entry_date
|
|
# reaction_close: enter at close of reaction_date itself
|
|
# next_open_after_reaction_close: enter at open of next trading day
|
|
if entry_convention == "reaction_close":
|
|
entry_date = reaction_date
|
|
else:
|
|
entry_date = next_trading_day(reaction_date)
|
|
|
|
# 3. Fetch price bars (entry_date + _LOOK_AHEAD_DAYS trading days)
|
|
fetch_start = entry_date
|
|
trading_days = trading_days_between(entry_date, entry_date + dt.timedelta(days=20))
|
|
fetch_end = trading_days[_LOOK_AHEAD_DAYS] if len(trading_days) > _LOOK_AHEAD_DAYS else trading_days[-1]
|
|
|
|
try:
|
|
price_resp = await price_svc.get_daily_bars(
|
|
ticker=ticker,
|
|
start=fetch_start.isoformat(),
|
|
end=fetch_end.isoformat(),
|
|
)
|
|
raw_bars = [b.model_dump() for b in price_resp.bars]
|
|
except Exception as exc:
|
|
logger.warning("label_price_unavailable", ticker=ticker, error=str(exc))
|
|
return EventLabel(
|
|
event_id=event.event_id,
|
|
entry_convention=entry_convention,
|
|
reaction_date=reaction_date,
|
|
entry_date=None,
|
|
label_status="unavailable",
|
|
invalid_event_for_labeling=False,
|
|
label_version=LABEL_VERSION,
|
|
)
|
|
|
|
# Filter bars from entry_date onwards, sorted ascending
|
|
bars = sorted(
|
|
[b for b in raw_bars if b.get("date") and b["date"] >= entry_date.isoformat()],
|
|
key=lambda b: b["date"],
|
|
)
|
|
|
|
if not bars:
|
|
# For reaction_close, the bar may not exist yet if market hasn't closed.
|
|
# Return "pending" so the pipeline can regenerate after market close.
|
|
if entry_convention == "reaction_close" and entry_date >= dt.date.today():
|
|
label_status_no_bar = "pending"
|
|
else:
|
|
label_status_no_bar = "unavailable"
|
|
return EventLabel(
|
|
event_id=event.event_id,
|
|
entry_convention=entry_convention,
|
|
reaction_date=reaction_date,
|
|
entry_date=entry_date,
|
|
label_status=label_status_no_bar,
|
|
invalid_event_for_labeling=False,
|
|
label_version=LABEL_VERSION,
|
|
)
|
|
|
|
# 4. Determine entry_price
|
|
first_bar = bars[0]
|
|
if entry_convention == "next_open_after_reaction_close":
|
|
entry_price = _safe_decimal(first_bar.get("open"))
|
|
else:
|
|
entry_price = _safe_decimal(first_bar.get("close"))
|
|
|
|
if entry_price is None or entry_price == 0:
|
|
return EventLabel(
|
|
event_id=event.event_id,
|
|
entry_convention=entry_convention,
|
|
reaction_date=reaction_date,
|
|
entry_date=entry_date,
|
|
entry_price=None,
|
|
label_status="unavailable",
|
|
invalid_event_for_labeling=True,
|
|
label_version=LABEL_VERSION,
|
|
)
|
|
|
|
# 5. Forward bars (exclude entry bar itself for 1D/3D/5D)
|
|
forward_bars = bars[1:] # Day 1+ after entry
|
|
|
|
label_status = "ok"
|
|
if len(forward_bars) < 20:
|
|
label_status = "truncated"
|
|
|
|
# 1D return
|
|
fwd_1d = _pct_return(entry_price, _safe_decimal(bars[1]["close"])) if len(bars) > 1 else None
|
|
|
|
# 3D labels
|
|
lbl_3d = _compute_labels_from_bars(entry_price, forward_bars, 3)
|
|
# 5D labels
|
|
lbl_5d = _compute_labels_from_bars(entry_price, forward_bars, 5)
|
|
# 10D labels
|
|
lbl_10d = _compute_labels_from_bars(entry_price, forward_bars, 10)
|
|
# 20D labels
|
|
lbl_20d = _compute_labels_from_bars(entry_price, forward_bars, 20)
|
|
|
|
return EventLabel(
|
|
event_id=event.event_id,
|
|
entry_convention=entry_convention,
|
|
reaction_date=reaction_date,
|
|
entry_date=entry_date,
|
|
entry_price=entry_price,
|
|
fwd_return_1d=fwd_1d,
|
|
fwd_return_3d=lbl_3d.get("fwd_return"),
|
|
fwd_return_5d=lbl_5d.get("fwd_return"),
|
|
fwd_return_10d=lbl_10d.get("fwd_return"),
|
|
fwd_return_20d=lbl_20d.get("fwd_return"),
|
|
hit_pos_1r_within_3d=lbl_3d.get("hit_pos_1r"),
|
|
hit_neg_1r_within_3d=lbl_3d.get("hit_neg_1r"),
|
|
close_up_after_3d=lbl_3d.get("close_up"),
|
|
close_up_after_5d=lbl_5d.get("close_up"),
|
|
mfe_3d=lbl_3d.get("mfe"),
|
|
mae_3d=lbl_3d.get("mae"),
|
|
mfe_5d=lbl_5d.get("mfe"),
|
|
mae_5d=lbl_5d.get("mae"),
|
|
mfe_10d=lbl_10d.get("mfe"),
|
|
mae_10d=lbl_10d.get("mae"),
|
|
mfe_20d=lbl_20d.get("mfe"),
|
|
mae_20d=lbl_20d.get("mae"),
|
|
bars_to_mfe_3d=lbl_3d.get("bars_to_mfe"),
|
|
bars_to_mae_3d=None,
|
|
days_to_peak_close_5d=lbl_5d.get("days_to_peak_close"),
|
|
label_status=label_status,
|
|
invalid_event_for_labeling=False,
|
|
label_version=LABEL_VERSION,
|
|
)
|