67 lines
2.5 KiB
Python
67 lines
2.5 KiB
Python
"""Minimal persistence of computed SP features into the lake features parquet.
|
|
|
|
Flow: compute sp_* features per symbol (examples/sp_features.py) and MERGE them
|
|
into features/market=US/timeframe=1d/symbol=*.parquet so TACHandler /
|
|
LakeFeatureProvider can route `$sp_ou_theta` etc. from the workflow YAML.
|
|
|
|
Run after backfilling bars; re-run drops stale sp_* columns first (see note).
|
|
|
|
python examples/persist_sp_features.py --market US --timeframe 1d
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import os
|
|
|
|
import pandas as pd
|
|
|
|
from tac_qlib.data.config import LakeConfig, NON_FEATURE_COLUMNS
|
|
from examples.sp_features import build_sp_features
|
|
|
|
#: columns owned by this feature family (replaced on re-runs, never duplicated)
|
|
SP_PREFIX = "sp_"
|
|
|
|
|
|
def persist_symbol(lake: LakeConfig, timeframe: str, symbol: str) -> None:
|
|
bars_path = lake.bar_path(timeframe, symbol)
|
|
feats_path = lake.features_path(timeframe, symbol)
|
|
if not bars_path.exists():
|
|
return
|
|
bars = pd.read_parquet(bars_path)
|
|
feats = build_sp_features(bars)
|
|
# bars have a single 't'/'date' column; align feature rows to it
|
|
feats = feats.drop(columns=[c for c in NON_FEATURE_COLUMNS if c in feats.columns], errors="ignore")
|
|
|
|
feats_path.parent.mkdir(parents=True, exist_ok=True)
|
|
if feats_path.exists():
|
|
existing = pd.read_parquet(feats_path)
|
|
# drop stale sp_* columns before merging (idempotent re-runs)
|
|
existing = existing[[c for c in existing.columns if not c.startswith(SP_PREFIX)]]
|
|
merged = pd.merge(existing, feats, on="t", how="left", suffixes=("", "_dup"))
|
|
merged = merged.loc[:, ~merged.columns.str.endswith("_dup")]
|
|
# keep original column order + new sp_* appended
|
|
merged.to_parquet(feats_path, index=False)
|
|
else:
|
|
feats.to_parquet(feats_path, index=False)
|
|
print(f"persisted {symbol}: {len(feats.columns) - 1} sp_* features")
|
|
|
|
|
|
def main() -> None:
|
|
ap = argparse.ArgumentParser()
|
|
ap.add_argument("--market", default="US")
|
|
ap.add_argument("--timeframe", default="1d")
|
|
ap.add_argument("--symbols", default="", help="comma-separated; default: all lake symbols")
|
|
args = ap.parse_args()
|
|
lake_root = os.environ.get("TAC_LAKE_DIR")
|
|
if not lake_root:
|
|
raise SystemExit("TAC_LAKE_DIR is required")
|
|
lake = LakeConfig(lake_root, args.market)
|
|
symbols = [s.strip().upper() for s in args.symbols.split(",") if s.strip()] or lake.load_symbols()
|
|
for symbol in symbols:
|
|
persist_symbol(lake, args.timeframe, symbol)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|