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.

122 lines
4.0 KiB
Python

"""Dataset Export: snapshot features + labels to Parquet."""
from __future__ import annotations
import argparse
import asyncio
import json
import yaml
from libs.common.config import get_settings
from libs.common.logging import configure_logging, get_logger
from libs.db.session import get_session
from libs.export.snapshot_export import export_dataset_snapshot
logger = get_logger(__name__)
async def run_dataset_export(
snapshot_id: str | None,
split_policy: str,
output_dir: str,
feature_versions: list[str] | None = None,
label_version: str = "label-2.0.0",
symbols: list[str] | None = None,
universe_profile: str | None = None,
start_date: str | None = None,
end_date: str | None = None,
) -> dict:
import datetime as dt
async with get_session() as session:
manifest = await export_dataset_snapshot(
session=session,
snapshot_id=snapshot_id,
split_policy=split_policy,
output_dir=output_dir,
feature_versions=feature_versions,
label_version=label_version,
symbols=symbols,
universe_profile=universe_profile,
start_date=dt.date.fromisoformat(start_date) if start_date else None,
end_date=dt.date.fromisoformat(end_date) if end_date else None,
)
return manifest
def main() -> None:
parser = argparse.ArgumentParser(description="Dataset Export")
parser.add_argument("--snapshot-id", default=None, help="Snapshot ID (UUID, auto-generated if omitted)")
parser.add_argument(
"--split-policy",
default="temporal_70_15_15",
help="Split policy string (default: temporal_70_15_15)",
)
parser.add_argument(
"--output-dir",
default="./data/datasets/snapshots",
help="Output directory for Parquet files",
)
parser.add_argument(
"--feature-versions",
nargs="+",
default=None,
help="Feature versions to merge (e.g. market_v1 event_v1 financial_v1)",
)
parser.add_argument(
"--label-version",
default="label-2.0.0",
help="Label version to export (default: label-2.0.0)",
)
parser.add_argument(
"--symbols-file",
default=None,
help="YAML file with 'symbols' list to filter export (e.g. configs/symbols_midcap.yaml)",
)
parser.add_argument(
"--universe-profile",
default=None,
help="Named live screener profile to filter export (e.g. midlarge-liquid-long-v1)",
)
parser.add_argument("--start-date", default=None, metavar="YYYY-MM-DD", help="Inclusive event_date lower bound")
parser.add_argument("--end-date", default=None, metavar="YYYY-MM-DD", help="Inclusive event_date upper bound")
parser.add_argument("--json", action="store_true", help="Print manifest JSON to stdout")
args = parser.parse_args()
settings = get_settings()
configure_logging(settings.log_level)
symbols = None
if args.symbols_file:
with open(args.symbols_file) as f:
cfg = yaml.safe_load(f) or {}
symbols = [s for s in cfg.get("symbols", []) if isinstance(s, str)]
print(f"Filtering to {len(symbols)} symbols from {args.symbols_file}")
manifest = asyncio.run(
run_dataset_export(
snapshot_id=args.snapshot_id,
split_policy=args.split_policy,
output_dir=args.output_dir,
feature_versions=args.feature_versions,
label_version=args.label_version,
symbols=symbols,
universe_profile=args.universe_profile,
start_date=args.start_date,
end_date=args.end_date,
)
)
if args.json:
print(json.dumps(manifest, indent=2))
else:
print(f"Snapshot exported: {manifest['snapshot_id']}")
print(f" Output: {manifest['output_dir']}")
print(f" Total rows: {manifest['total_rows']}")
for split, count in manifest["row_counts"].items():
print(f" {split}: {count} rows")
if __name__ == "__main__":
main()