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.
fithia2/scripts/enrich_catalyst_persistence...

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()