prepare_input.py 1.9 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061
  1. #!/usr/bin/env python3
  2. """Normalize arbitrary MID and anchor-date inputs for the risk feature pipeline."""
  3. from __future__ import annotations
  4. import argparse
  5. import sys
  6. from pathlib import Path
  7. import pandas as pd
  8. PROJECT_DIR = Path(__file__).resolve().parents[1]
  9. if str(PROJECT_DIR) not in sys.path:
  10. sys.path.insert(0, str(PROJECT_DIR))
  11. from src.user_input import load_users, normalize_users # noqa: E402
  12. def parse_user(value: str) -> dict[str, str]:
  13. parts = value.split(":", 3)
  14. if len(parts) < 2:
  15. raise ValueError("--user格式应为 MID:yyyyMMdd[:用户类型[:子分类]]")
  16. return {
  17. "mid": parts[0].strip(),
  18. "anchor_dt": parts[1].strip(),
  19. "user_type": parts[2].strip() if len(parts) > 2 else "待分类",
  20. "sub_category": parts[3].strip() if len(parts) > 3 else "待分类",
  21. }
  22. def main() -> None:
  23. parser = argparse.ArgumentParser(description=__doc__)
  24. parser.add_argument("--input", type=Path, help="已有CSV,至少包含MID和锚点日期")
  25. parser.add_argument(
  26. "--user",
  27. action="append",
  28. default=[],
  29. help="可重复传入:MID:yyyyMMdd[:用户类型[:子分类]]",
  30. )
  31. parser.add_argument("--lookback-days", type=int, default=180)
  32. parser.add_argument("--output", type=Path, required=True)
  33. args = parser.parse_args()
  34. if bool(args.input) == bool(args.user):
  35. parser.error("--input和--user必须且只能使用一种")
  36. users = (
  37. load_users(args.input, args.lookback_days)
  38. if args.input
  39. else normalize_users(
  40. pd.DataFrame([parse_user(value) for value in args.user]),
  41. lookback_days=args.lookback_days,
  42. )
  43. )
  44. args.output.parent.mkdir(parents=True, exist_ok=True)
  45. users.to_csv(args.output, index=False)
  46. print(f"users={len(users)} unique_mids={users['mid'].nunique()} output={args.output}")
  47. if __name__ == "__main__":
  48. main()