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.
163 lines
5.5 KiB
Python
163 lines
5.5 KiB
Python
"""Enrich snapshot with PIT-safe trailing catalyst persistence features.
|
|
|
|
This derives simple trailing event-density features from the snapshot rows
|
|
themselves, using only events that happened strictly before the current row.
|
|
All splits are loaded together so valid/test rows can safely use earlier
|
|
train-period events.
|
|
|
|
Usage:
|
|
uv run python3 scripts/enrich_catalyst_persistence_features.py \
|
|
--input data/parquet/midlarge-liquid-long-v1_bucketfix_full_audit_canonical \
|
|
--output data/parquet/midlarge-liquid-long-v1_bucketfix_full_audit_canonical_cp
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import datetime as dt
|
|
import json
|
|
from collections import defaultdict, deque
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
import pandas as pd
|
|
import pyarrow as pa
|
|
import pyarrow.parquet as pq
|
|
|
|
|
|
LOOKBACK_WINDOWS = (20, 60)
|
|
PERSISTENCE_FEATURES = [
|
|
"prior_catalyst_count_20d",
|
|
"prior_catalyst_count_60d",
|
|
"prior_catalyst_type_diversity_20d",
|
|
"prior_catalyst_type_diversity_60d",
|
|
]
|
|
|
|
|
|
def _parse_date(value: Any) -> dt.date | None:
|
|
if value is None:
|
|
return None
|
|
if isinstance(value, dt.datetime):
|
|
return value.date()
|
|
if isinstance(value, dt.date):
|
|
return value
|
|
text = str(value).strip()
|
|
if not text:
|
|
return None
|
|
try:
|
|
return dt.date.fromisoformat(text[:10])
|
|
except ValueError:
|
|
return None
|
|
|
|
|
|
def _load_all_rows(input_dir: Path) -> pd.DataFrame:
|
|
frames: list[pd.DataFrame] = []
|
|
split_order = {"train": 0, "valid": 1, "test": 2}
|
|
for split in ("train", "valid", "test"):
|
|
parquet_path = input_dir / f"{split}.parquet"
|
|
if not parquet_path.exists():
|
|
continue
|
|
df = pq.read_table(parquet_path).to_pandas()
|
|
df["__split_name"] = split
|
|
df["__split_order"] = split_order[split]
|
|
df["__split_row_idx"] = range(len(df))
|
|
frames.append(df)
|
|
if not frames:
|
|
return pd.DataFrame()
|
|
combined = pd.concat(frames, ignore_index=True)
|
|
combined["__event_date_obj"] = combined["event_date"].map(_parse_date)
|
|
combined = combined.sort_values(
|
|
by=["ticker", "__event_date_obj", "__split_order", "__split_row_idx"],
|
|
kind="stable",
|
|
na_position="last",
|
|
).reset_index(drop=True)
|
|
return combined
|
|
|
|
|
|
def compute_catalyst_persistence_features(combined: pd.DataFrame) -> pd.DataFrame:
|
|
if combined.empty:
|
|
out = combined.copy()
|
|
for feature in PERSISTENCE_FEATURES:
|
|
out[feature] = pd.Series(dtype="float64")
|
|
return out
|
|
|
|
queues: dict[int, dict[str, deque[tuple[dt.date, str]]]] = {
|
|
window: defaultdict(deque) for window in LOOKBACK_WINDOWS
|
|
}
|
|
feature_values = {feature: [None] * len(combined) for feature in PERSISTENCE_FEATURES}
|
|
|
|
for idx, row in combined.iterrows():
|
|
ticker = str(row.get("ticker") or "").upper()
|
|
event_date = row.get("__event_date_obj")
|
|
event_type = str(row.get("event_type") or "UNKNOWN") or "UNKNOWN"
|
|
|
|
if not ticker or event_date is None:
|
|
continue
|
|
|
|
for window in LOOKBACK_WINDOWS:
|
|
dq = queues[window][ticker]
|
|
cutoff = event_date - dt.timedelta(days=window)
|
|
while dq and dq[0][0] < cutoff:
|
|
dq.popleft()
|
|
|
|
feature_values[f"prior_catalyst_count_{window}d"][idx] = float(len(dq))
|
|
feature_values[f"prior_catalyst_type_diversity_{window}d"][idx] = float(
|
|
len({prior_type for _, prior_type in dq})
|
|
)
|
|
|
|
for window in LOOKBACK_WINDOWS:
|
|
queues[window][ticker].append((event_date, event_type))
|
|
|
|
out = combined.copy()
|
|
for feature, values in feature_values.items():
|
|
out[feature] = values
|
|
return out
|
|
|
|
|
|
def enrich_snapshot_dir(input_dir: Path, output_dir: Path) -> None:
|
|
output_dir.mkdir(parents=True, exist_ok=True)
|
|
|
|
combined = _load_all_rows(input_dir)
|
|
if combined.empty:
|
|
manifest_src = input_dir / "manifest.json"
|
|
if manifest_src.exists():
|
|
output_dir.joinpath("manifest.json").write_text(manifest_src.read_text())
|
|
return
|
|
|
|
enriched = compute_catalyst_persistence_features(combined)
|
|
|
|
manifest_src = input_dir / "manifest.json"
|
|
if manifest_src.exists():
|
|
manifest = json.loads(manifest_src.read_text())
|
|
source_snapshot_id = manifest.get("snapshot_id")
|
|
manifest["snapshot_id"] = output_dir.name
|
|
manifest["output_dir"] = str(output_dir.resolve())
|
|
manifest["created_at_utc"] = dt.datetime.now(dt.UTC).isoformat()
|
|
manifest["enrichment_source_snapshot_id"] = source_snapshot_id
|
|
manifest["export_enrichments"] = sorted(
|
|
set((manifest.get("export_enrichments") or []) + PERSISTENCE_FEATURES)
|
|
)
|
|
output_dir.joinpath("manifest.json").write_text(json.dumps(manifest, indent=2))
|
|
|
|
for split in ("train", "valid", "test"):
|
|
split_df = enriched.loc[enriched["__split_name"] == split].copy()
|
|
if split_df.empty:
|
|
continue
|
|
split_df = split_df.drop(
|
|
columns=["__split_name", "__split_order", "__split_row_idx", "__event_date_obj"],
|
|
errors="ignore",
|
|
)
|
|
split_table = pa.Table.from_pandas(split_df, preserve_index=False)
|
|
pq.write_table(split_table, output_dir / f"{split}.parquet")
|
|
|
|
|
|
def main() -> None:
|
|
parser = argparse.ArgumentParser()
|
|
parser.add_argument("--input", required=True)
|
|
parser.add_argument("--output", required=True)
|
|
args = parser.parse_args()
|
|
enrich_snapshot_dir(Path(args.input), Path(args.output))
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|