ringguard / scoring /features.py
Rohit Patil
RingGuard Phase 3 - initial HF Spaces deploy
753201e
Raw History Blame Contribute Delete
9.85 kB
"""
features.py
-----------
Feature extraction for RingGuard risk model.
Given a candidate cluster (from candidate_clusters.json) plus the matching
orders.json, customers.json, and optionally graph data, produces exactly 13
features per cluster as specified in the Phase 2 plan.
Feature set:
size, strong_link_count, weak_link_count, strong_link_density,
has_same_contact_edge, order_count, total_amount, avg_order_amount,
return_rate, chargeback_rate, shared_coupon_count, time_span_days,
order_velocity, formation_type_strong
Usage:
from scoring.features import extract_features, label
features = extract_features(cluster, orders, graph=G)
lbl = label(cluster, ground_truth)
"""
import json
import os
import pickle
from datetime import datetime
from typing import Optional
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _load_json(path: str):
with open(path, encoding="utf-8") as f:
return json.load(f)
def _load_graph(path: str):
try:
import networkx as nx
return nx.read_gpickle(path)
except AttributeError:
with open(path, "rb") as f:
return pickle.load(f)
# ---------------------------------------------------------------------------
# Feature extraction
# ---------------------------------------------------------------------------
def extract_features(
cluster: dict,
orders: list[dict],
graph=None,
) -> dict:
"""
Extract exactly 13 features from a candidate cluster.
Parameters
----------
cluster : dict
One entry from candidate_clusters.json. Must have:
customer_ids, size, strong_link_count, weak_link_count, formation_type
orders : list[dict]
All orders from the dataset (orders.json). Filtered to cluster members.
graph : networkx.Graph, optional
The heterogeneous graph (graph.gpickle). Used to detect same_contact edges.
If None, has_same_contact_edge is computed from cluster data only
(will be 0 unless the cluster was formed via same_contact; this is
a reasonable fallback since same_contact rings are always strong_anchored).
Returns
-------
dict with exactly 13 feature keys + cluster_id for identification.
"""
member_set = set(cluster["customer_ids"])
size = cluster["size"]
strong_link_count = cluster["strong_link_count"]
weak_link_count = cluster["weak_link_count"]
formation_type = cluster.get("formation_type", "strong_anchored")
# Feature 1: size
# Feature 2: strong_link_count (from cluster record)
# Feature 3: weak_link_count (from cluster record)
# Feature 4: strong_link_density
max_possible = max(1, size * (size - 1) / 2)
strong_link_density = strong_link_count / max_possible
# Feature 5: has_same_contact_edge
# Check graph for same_contact edges between cluster members
has_same_contact_edge = 0
if graph is not None:
for u, v, data in graph.edges(data=True):
if data.get("relation") == "same_contact":
if u in member_set and v in member_set:
has_same_contact_edge = 1
break
# Filter orders to cluster members
cluster_orders = [o for o in orders if o["customer_id"] in member_set]
# Feature 6: order_count
order_count = len(cluster_orders)
# Feature 7: total_amount
total_amount = sum(o.get("amount", 0) for o in cluster_orders)
# Feature 8: avg_order_amount
avg_order_amount = total_amount / max(1, order_count)
# Feature 9: return_rate
returned = sum(1 for o in cluster_orders if o.get("status") == "returned")
return_rate = returned / max(1, order_count)
# Feature 10: chargeback_rate
chargebacks = sum(1 for o in cluster_orders if o.get("status") == "chargeback")
chargeback_rate = chargebacks / max(1, order_count)
# Feature 11: shared_coupon_count
# Count distinct coupon_id values used by 2+ different cluster members
coupon_to_members: dict[str, set] = {}
for o in cluster_orders:
cid_val = o.get("coupon_id")
if cid_val is not None:
if cid_val not in coupon_to_members:
coupon_to_members[cid_val] = set()
coupon_to_members[cid_val].add(o["customer_id"])
shared_coupon_count = sum(
1 for members in coupon_to_members.values() if len(members) >= 2
)
# Features 12 + 13: time_span_days and order_velocity
timestamps = []
for o in cluster_orders:
ts_str = o.get("timestamp", "")
if ts_str:
try:
timestamps.append(datetime.fromisoformat(ts_str))
except ValueError:
pass
if len(timestamps) >= 2:
time_span_days = (max(timestamps) - min(timestamps)).total_seconds() / 86400.0
else:
time_span_days = 0.0
order_velocity = order_count / max(1.0, time_span_days)
# Feature 14 (index 13): formation_type_strong
formation_type_strong = 1 if formation_type == "strong_anchored" else 0
return {
"cluster_id": cluster["cluster_id"],
"size": size,
"strong_link_count": strong_link_count,
"weak_link_count": weak_link_count,
"strong_link_density": strong_link_density,
"has_same_contact_edge": has_same_contact_edge,
"order_count": order_count,
"total_amount": total_amount,
"avg_order_amount": avg_order_amount,
"return_rate": return_rate,
"chargeback_rate": chargeback_rate,
"shared_coupon_count": shared_coupon_count,
"time_span_days": time_span_days,
"order_velocity": order_velocity,
"formation_type_strong": formation_type_strong,
}
# ---------------------------------------------------------------------------
# Labeling
# ---------------------------------------------------------------------------
FEATURE_COLUMNS = [
"size",
"strong_link_count",
"weak_link_count",
"strong_link_density",
"has_same_contact_edge",
"order_count",
"total_amount",
"avg_order_amount",
"return_rate",
"chargeback_rate",
"shared_coupon_count",
"time_span_days",
"order_velocity",
"formation_type_strong",
]
def label(cluster: dict, ground_truth: dict) -> int:
"""
Returns 1 if >= 50% of the cluster's customer_ids overlap with any
ground-truth ring, else 0.
Parameters
----------
cluster : dict
One entry from candidate_clusters.json.
ground_truth : dict
Loaded ground_truth.json. Must have a "rings" key.
"""
cluster_members = set(cluster["customer_ids"])
threshold = len(cluster_members) * 0.5
for ring in ground_truth.get("rings", []):
ring_members = set(ring["customer_ids"])
overlap = len(cluster_members & ring_members)
if overlap >= threshold:
return 1
return 0
# ---------------------------------------------------------------------------
# Batch extraction
# ---------------------------------------------------------------------------
def extract_all(
clusters: list[dict],
orders: list[dict],
ground_truth: dict = None,
graph=None,
) -> list[dict]:
"""
Extract features for every cluster in the list.
If ground_truth is provided, adds a 'label' field to each row.
Returns list of feature dicts (one per cluster), ready for pandas or json.
"""
rows = []
for cluster in clusters:
row = extract_features(cluster, orders, graph=graph)
if ground_truth is not None:
row["label"] = label(cluster, ground_truth)
rows.append(row)
return rows
# ---------------------------------------------------------------------------
# CLI convenience: run directly to verify feature extraction on a dataset
# ---------------------------------------------------------------------------
if __name__ == "__main__":
import argparse
import sys
parser = argparse.ArgumentParser()
parser.add_argument("--data-dir", default=None,
help="Path to data directory. Defaults to data/synthetic_train/")
args = parser.parse_args()
base = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
data_dir = args.data_dir or os.path.join(base, "data", "synthetic_train")
clusters_path = os.path.join(data_dir, "candidate_clusters.json")
orders_path = os.path.join(data_dir, "orders.json")
gt_path = os.path.join(data_dir, "ground_truth.json")
graph_path = os.path.join(data_dir, "graph.gpickle")
clusters = _load_json(clusters_path)
orders = _load_json(orders_path)
ground_truth = _load_json(gt_path)
graph = _load_graph(graph_path) if os.path.exists(graph_path) else None
rows = extract_all(clusters, orders, ground_truth=ground_truth, graph=graph)
# Sanity checks
assert all(len(r) == len(FEATURE_COLUMNS) + 2 for r in rows), \
"Row has wrong number of features (expected 14 features + cluster_id + label)"
null_count = sum(
1 for r in rows for k in FEATURE_COLUMNS if r.get(k) is None
)
assert null_count == 0, f"Null values found in feature table: {null_count}"
labels = [r["label"] for r in rows]
pos = sum(labels)
neg = len(labels) - pos
print(f"Feature extraction OK: {len(rows)} clusters")
print(f" Positive labels (fraud ring): {pos}")
print(f" Negative labels (benign): {neg}")
print(f" Null values in features: {null_count}")
print("\nSample row:")
sample = rows[0]
for k in FEATURE_COLUMNS:
print(f" {k:30s}: {sample[k]:.4f}" if isinstance(sample[k], float)
else f" {k:30s}: {sample[k]}")