"""Patch ftb_fix_v2 snapshot by merging fresh base export with existing canonical enrichments. Strategy: - 00_base has correct market features + labels (built from freshly-corrected DB) - Existing canonical has enriched features (tier2/tier3/technical/macro/prior_drift) - These enrichments don't depend on reaction_date → valid for ALL rows including 1,153 affected - For rows in both: use base features, copy enrichments from canonical - For new rows (not in canonical): use base features, NULL enrichments """ from __future__ import annotations import json import datetime as dt from pathlib import Path import pyarrow as pa import pyarrow.parquet as pq CANONICAL_DIR = Path("data/parquet/midlarge-liquid-long-v1_bucketfix_full_audit_canonical") BASE_DIR = Path("data/parquet/.midlarge-liquid-long-v1_bucketfix_full_audit_canonical_ftb_fix_v2.stage.3eb6jxkd/00_base") OUTPUT_DIR = Path("data/parquet/midlarge-liquid-long-v1_bucketfix_full_audit_canonical_ftb_fix_v2") # Columns that come FROM the base export (fresh correct data from DB) # These are the columns affected by the filing_time_bucket fix BASE_COLUMNS = { # Market features (reaction-date-dependent) "reaction_day_return", "volume_ratio_20d", "avg_dollar_volume_20d", "gap_size", "close_location", "event_close", "event_volume", "reaction_day_low", "reaction_day_high", # Pre-event features (shift by 1 day for affected rows, minor change - use fresh) "pre_event_momentum_20d", "price_vs_sma20", "pre_event_volatility_20d", "pre_event_rsi_14", "pre_event_bb_position", "pre_event_obv_slope_20d", "atr_14", # Event features "filing_time_bucket", # Labels (reaction_date/entry_date/entry_price/fwd_returns - all depend on fix) "reaction_date", "entry_date", "entry_price", "fwd_return_1d", "fwd_return_3d", "fwd_return_5d", "fwd_return_10d", "fwd_return_20d", "mfe_3d", "mae_3d", "mfe_5d", "mae_5d", "mfe_10d", "mae_10d", "mfe_20d", "mae_20d", "hit_pos_1r_within_3d", "hit_neg_1r_within_3d", "close_up_after_3d", "close_up_after_5d", "bars_to_mfe_3d", "days_to_peak_close_5d", "label_status", "label_version", # Metadata "event_date", "market_cap_proxy", "exchange_proxy", } def patch_split(split: str) -> int: base_path = BASE_DIR / f"{split}.parquet" canonical_path = CANONICAL_DIR / f"{split}.parquet" output_path = OUTPUT_DIR / f"{split}.parquet" if not base_path.exists(): print(f" No {split} split in base, skipping") return 0 base_table = pq.read_table(base_path) n = len(base_table) print(f" Base {split}: {n} rows, {len(base_table.column_names)} cols") # Build index from existing canonical by event_id canonical_data: dict[str, dict] = {} if canonical_path.exists(): canon_table = pq.read_table(canonical_path) canon_cols = set(canon_table.column_names) print(f" Canonical {split}: {len(canon_table)} rows, {len(canon_cols)} cols") # Get event_id column name (might be 'event_id' or other) id_col = "event_id" if "event_id" in canon_cols else None if id_col: ids = canon_table.column(id_col).to_pylist() # Build dict: event_id -> row dict for enrichment columns enrichment_cols = [c for c in canon_cols if c not in BASE_COLUMNS and c != id_col] for j, eid in enumerate(ids): row = {} for col in enrichment_cols: val = canon_table.column(col)[j].as_py() row[col] = val canonical_data[eid] = row print(f" Indexed {len(canonical_data)} canonical rows") # Build output table # Start with all base columns base_cols = base_table.column_names # Find enrichment columns from canonical not in base all_canonical_cols = set() if canonical_path.exists(): canon_table2 = pq.read_table(canonical_path) all_canonical_cols = set(canon_table2.column_names) enrich_cols = sorted([c for c in all_canonical_cols if c not in base_cols and c not in {"event_id"}]) print(f" Adding {len(enrich_cols)} enrichment columns from canonical") # Get event_ids from base base_id_col = "event_id" if "event_id" in base_cols else None base_ids = base_table.column(base_id_col).to_pylist() if base_id_col else [None] * n # For each enrichment column, build array result_table = base_table for col in enrich_cols: # Get type from canonical if canonical_path.exists(): canon_field = canon_table2.schema.field(col) col_type = canon_field.type else: col_type = pa.float64() values = [] for eid in base_ids: val = canonical_data.get(eid, {}).get(col) if eid else None values.append(val) try: arr = pa.array(values, type=col_type) except Exception: arr = pa.array(values) result_table = result_table.append_column(col, arr) pq.write_table(result_table, output_path) print(f" Written {len(result_table)} rows, {len(result_table.column_names)} cols to {split}.parquet") return n def main(): OUTPUT_DIR.mkdir(parents=True, exist_ok=True) print("Patching ftb_fix_v2 snapshot...") total = 0 for split in ["train", "valid", "test"]: print(f"\nProcessing {split}:") total += patch_split(split) # Copy manifest from base, update metadata base_manifest = json.loads((BASE_DIR / "manifest.json").read_text()) base_manifest["snapshot_id"] = OUTPUT_DIR.name base_manifest["output_dir"] = str(OUTPUT_DIR.resolve()) base_manifest["created_at_utc"] = dt.datetime.now(dt.UTC).isoformat() base_manifest["ftb_fix_note"] = "Patched: base features from corrected DB, enrichments from existing canonical" (OUTPUT_DIR / "manifest.json").write_text(json.dumps(base_manifest, indent=2)) print(f"\nDone! Total rows: {total}") print(f"Output: {OUTPUT_DIR}") if __name__ == "__main__": main()