run_social_features.py 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259
  1. #!/usr/bin/env python3
  2. """Query and aggregate share relationships and source-group features."""
  3. from __future__ import annotations
  4. import argparse
  5. import json
  6. import sys
  7. from datetime import datetime, timedelta
  8. from pathlib import Path
  9. import pandas as pd
  10. PROJECT_DIR = Path(__file__).resolve().parents[1]
  11. WORKSPACE_ROOT = Path(__file__).resolve().parents[2]
  12. for path in (PROJECT_DIR, WORKSPACE_ROOT):
  13. if str(path) not in sys.path:
  14. sys.path.insert(0, str(path))
  15. from odps_module import ODPSClient # noqa: E402
  16. from src.social_logic import ( # noqa: E402
  17. INVALID_GROUP_IDS,
  18. aggregate_groups,
  19. aggregate_relationships,
  20. build_group_sql,
  21. build_relationship_sql,
  22. latest_partition,
  23. )
  24. from src.user_input import load_users # noqa: E402
  25. RELATIONSHIP_COLUMNS = [
  26. "sample_mid", "user_type", "sub_category", "anchor_dt",
  27. "window_start_dt", "window_end_dt", "click_event_id", "click_dt",
  28. "click_ts", "shareid", "clickobjectid", "hotsencetype", "valid_opengid",
  29. "sessionid", "subsessionid", "share_dt", "share_ts", "source_share_mid",
  30. "is_self_source", "source_is_blacklist", "source_risk_type_list",
  31. "share_click_gap_seconds", "previous_clickobjectid", "previous_click_ts",
  32. "is_consecutive_same_card",
  33. ]
  34. GROUP_COLUMNS = [
  35. "sample_mid", "user_type", "sub_category", "window_start_dt",
  36. "window_end_dt", "opengid", "group_path_pv", "group_visit_uv",
  37. "group_risk_user_uv", "group_risk_user_rate", "group_source_sharer_uv",
  38. "group_risk_source_sharer_uv", "group_risk_source_sharer_rate",
  39. "group_real_play_pv", "group_real_play_uv", "group_real_play_user_rate",
  40. "group_real_play_per_visitor",
  41. ]
  42. def share_search_start(value: str) -> str:
  43. date = datetime.strptime(value, "%Y%m%d") - timedelta(days=180)
  44. return date.strftime("%Y%m%d")
  45. def build_group_input(users: pd.DataFrame, detail: pd.DataFrame) -> pd.DataFrame:
  46. required = {"用户id", "来源群id"}
  47. missing = required - set(detail.columns)
  48. if missing:
  49. raise ValueError("基础明细缺少字段: " + ", ".join(sorted(missing)))
  50. windows = users.set_index("mid")
  51. rows = []
  52. for record in detail[["用户id", "来源群id"]].to_dict("records"):
  53. mid = str(record["用户id"])
  54. if mid not in windows.index:
  55. raise ValueError(f"基础明细用户不在输入人员表: {mid}")
  56. group_ids = {
  57. value.strip()
  58. for value in str(record["来源群id"] or "").split(",")
  59. if value.strip() not in INVALID_GROUP_IDS
  60. }
  61. for opengid in sorted(group_ids):
  62. rows.append(
  63. {
  64. "sample_mid": mid,
  65. "user_type": windows.at[mid, "user_type"],
  66. "sub_category": windows.at[mid, "sub_category"],
  67. "window_start_dt": windows.at[mid, "window_start_dt"],
  68. "window_end_dt": windows.at[mid, "window_end_dt"],
  69. "opengid": opengid,
  70. }
  71. )
  72. return pd.DataFrame(
  73. rows,
  74. columns=[
  75. "sample_mid", "user_type", "sub_category", "window_start_dt",
  76. "window_end_dt", "opengid",
  77. ],
  78. ).drop_duplicates(["sample_mid", "opengid"])
  79. def read_instance(instance: object) -> pd.DataFrame:
  80. instance.wait_for_success()
  81. with instance.open_reader(tunnel=True) as reader:
  82. frame = reader.to_pandas()
  83. frame.columns = [str(column).lower() for column in frame.columns]
  84. return frame
  85. def write_outputs(
  86. users: pd.DataFrame,
  87. behavior_detail: pd.DataFrame,
  88. relationships: pd.DataFrame,
  89. group_detail: pd.DataFrame,
  90. output_dir: Path,
  91. ) -> None:
  92. user_relationship = aggregate_relationships(users, relationships)
  93. user_group = aggregate_groups(users, group_detail)
  94. social = user_relationship.merge(user_group, on="用户id", how="left", validate="one_to_one")
  95. extended = behavior_detail.merge(social, on="用户id", how="left", validate="one_to_one")
  96. order = {mid: index for index, mid in enumerate(users["mid"])}
  97. for frame, column in ((social, "用户id"), (extended, "用户id")):
  98. frame["_input_order"] = frame[column].map(order)
  99. frame.sort_values("_input_order", kind="stable", inplace=True)
  100. frame.drop(columns="_input_order", inplace=True)
  101. output_dir.mkdir(parents=True, exist_ok=True)
  102. social.to_csv(output_dir / "user_social_features.csv", index=False)
  103. relationships.to_csv(output_dir / "relationship_detail.csv", index=False)
  104. group_detail.to_csv(output_dir / "group_detail.csv", index=False)
  105. extended.to_csv(output_dir / "user_detail_with_social.csv", index=False)
  106. with pd.ExcelWriter(output_dir / "social_features.xlsx", engine="openpyxl") as writer:
  107. extended.to_excel(writer, sheet_name="用户扩展明细", index=False)
  108. social.to_excel(writer, sheet_name="用户社交特征", index=False)
  109. group_detail.to_excel(writer, sheet_name="来源群明细", index=False)
  110. relationships.to_excel(writer, sheet_name="分享关系明细", index=False)
  111. expected = len(users)
  112. if len(social) != expected or social["用户id"].nunique() != expected:
  113. raise AssertionError("社交特征未保持输入用户粒度")
  114. if len(extended) != expected or extended["用户id"].nunique() != expected:
  115. raise AssertionError("扩展明细未保持输入用户粒度")
  116. if not group_detail.empty and group_detail["opengid"].isin(INVALID_GROUP_IDS).any():
  117. raise AssertionError("群级结果包含无效群ID")
  118. gaps = pd.to_numeric(relationships["share_click_gap_seconds"], errors="coerce").dropna()
  119. if (gaps < 0).any():
  120. raise AssertionError("存在负分享点击时间差")
  121. def main() -> None:
  122. parser = argparse.ArgumentParser(description=__doc__)
  123. parser.add_argument("--input", type=Path, required=True)
  124. parser.add_argument("--behavior-detail", type=Path, required=True)
  125. parser.add_argument("--run-name", required=True)
  126. parser.add_argument("--lookback-days", type=int, default=180)
  127. parser.add_argument("--blacklist-dt", default="")
  128. parser.add_argument("--sql-only", action="store_true")
  129. parser.add_argument("--finalize-only", action="store_true")
  130. parser.add_argument("--collect-existing", action="store_true")
  131. parser.add_argument("--modules", default="")
  132. args = parser.parse_args()
  133. users = load_users(args.input, lookback_days=args.lookback_days)
  134. users["share_search_start_dt"] = users["window_start_dt"].map(share_search_start)
  135. behavior_detail = pd.read_csv(args.behavior_detail, dtype=str).fillna("")
  136. if set(behavior_detail["用户id"]) != set(users["mid"]):
  137. raise AssertionError("基础明细与输入人员的MID集合不一致")
  138. group_input = build_group_input(users, behavior_detail)
  139. output_dir = PROJECT_DIR / "output" / args.run_name / "social"
  140. raw_dir = output_dir / "raw"
  141. sql_dir = PROJECT_DIR / "sql" / args.run_name / "social"
  142. for directory in (output_dir, raw_dir, sql_dir):
  143. directory.mkdir(parents=True, exist_ok=True)
  144. users.to_csv(output_dir / "input_users.csv", index=False)
  145. group_input.to_csv(output_dir / "input_user_groups.csv", index=False)
  146. if args.finalize_only:
  147. relationships = pd.read_csv(raw_dir / "relationships.csv")
  148. groups = (
  149. pd.read_csv(raw_dir / "groups.csv")
  150. if (raw_dir / "groups.csv").exists()
  151. else pd.DataFrame(columns=GROUP_COLUMNS)
  152. )
  153. write_outputs(users, behavior_detail, relationships, groups, output_dir)
  154. print(f"output={output_dir / 'social_features.xlsx'}")
  155. return
  156. client = ODPSClient(project="loghubods")
  157. blacklist_dt = args.blacklist_dt or latest_partition(client, "ad_report_user_blacklist")
  158. queries = {"relationships": build_relationship_sql(users, blacklist_dt)}
  159. if not group_input.empty:
  160. queries["groups"] = build_group_sql(group_input, blacklist_dt)
  161. for name, sql in queries.items():
  162. (sql_dir / f"{name}.sql").write_text(sql + "\n", encoding="utf-8")
  163. (output_dir / "query_metadata.json").write_text(
  164. json.dumps(
  165. {
  166. "blacklist_dt": blacklist_dt,
  167. "users": len(users),
  168. "user_group_edges": len(group_input),
  169. "distinct_groups": group_input["opengid"].nunique() if len(group_input) else 0,
  170. },
  171. ensure_ascii=False,
  172. indent=2,
  173. )
  174. + "\n",
  175. encoding="utf-8",
  176. )
  177. if args.sql_only:
  178. print(f"sql_dir={sql_dir}")
  179. return
  180. selected = [item.strip() for item in args.modules.split(",") if item.strip()]
  181. unknown = sorted(set(selected) - set(queries))
  182. if unknown:
  183. raise ValueError("未知模块: " + ", ".join(unknown))
  184. selected_queries = {name: queries[name] for name in selected} if selected else queries
  185. runs_path = output_dir / "odps_runs.json"
  186. if args.collect_existing:
  187. history = json.loads(runs_path.read_text(encoding="utf-8"))
  188. latest = {item["name"]: item for item in history}
  189. instances = {
  190. name: client.odps.get_instance(latest[name]["instance_id"])
  191. for name in selected_queries
  192. }
  193. else:
  194. instances = {}
  195. new_runs = []
  196. for name, sql in selected_queries.items():
  197. instance = client.odps.run_sql(sql, hints={"odps.sql.allow.cartesian": "true"})
  198. instances[name] = instance
  199. run = {
  200. "name": name,
  201. "instance_id": instance.id,
  202. "logview": instance.get_logview_address(),
  203. }
  204. new_runs.append(run)
  205. print(f"[{name}] instance={run['instance_id']}", flush=True)
  206. print(f"[{name}] logview={run['logview']}", flush=True)
  207. history = json.loads(runs_path.read_text(encoding="utf-8")) if runs_path.exists() else []
  208. runs_path.write_text(
  209. json.dumps([*history, *new_runs], ensure_ascii=False, indent=2) + "\n",
  210. encoding="utf-8",
  211. )
  212. results = {}
  213. for name, instance in instances.items():
  214. frame = read_instance(instance)
  215. frame.to_csv(raw_dir / f"{name}.csv", index=False)
  216. results[name] = frame
  217. print(f"[{name}] rows={len(frame)}", flush=True)
  218. for name in queries:
  219. if name not in results:
  220. path = raw_dir / f"{name}.csv"
  221. if not path.exists():
  222. raise ValueError(f"缺少模块结果: {path}")
  223. results[name] = pd.read_csv(path)
  224. groups = results.get("groups", pd.DataFrame(columns=GROUP_COLUMNS))
  225. write_outputs(users, behavior_detail, results["relationships"], groups, output_dir)
  226. print(f"output={output_dir / 'social_features.xlsx'}")
  227. if __name__ == "__main__":
  228. main()