control_builder.py

#!/usr/bin/env python3
"""
Control Builder - Creates matched control cohorts for causal inference
Supports: geo-holdout, synthetic control, propensity score matching
"""

import logging
import os
from typing import Dict, List, Tuple

import numpy as np
from google.cloud import bigquery

logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)

PROJECT_ID = os.getenv("GCP_PROJECT_ID")
DATASET_ID = os.getenv("BQ_DATASET", "causal_roi")

client = bigquery.Client(project=PROJECT_ID)


def create_geo_holdout(
    path_id: str, treatment_regions: List[str], test_ratio: float = 0.2
) -> Tuple[str, float]:
    """
    Geographic holdout: randomly assign regions to treatment/control.

    Returns:
        cohort_id: identifier for control cohort
        quality_score: match quality (0-1)
    """
    all_regions = get_available_regions()
    np.random.shuffle(all_regions)

    n_control = int(len(all_regions) * test_ratio)
    control_regions = all_regions[:n_control]

    cohort_id = f"control_geo_{hash(path_id) % 1000}"

    # Quality score: balance check (region size similarity)
    quality_score = compute_balance_score(treatment_regions, control_regions)

    # Store control definition
    query = f"""
    INSERT INTO `{PROJECT_ID}.{DATASET_ID}.control_cohorts` 
    (cohort_id, path_id, method, geo_region, quality_score)
    VALUES (
        '{cohort_id}',
        '{path_id}',
        'geo_holdout',
        '{",".join(control_regions)}',
        {quality_score}
    )
    """
    client.query(query).result()

    logger.info(
        f"✓ Created geo holdout for {path_id}: {cohort_id} (quality={quality_score:.2f})"
    )
    return cohort_id, quality_score


def create_synthetic_control(
    path_id: str, treatment_features: Dict[str, float], donor_pool: List[str]
) -> Tuple[str, Dict[str, float], float]:
    """
    Synthetic control: weighted combination of untreated units.
    Uses convex optimization to match pre-treatment trends.

    Returns:
        cohort_id: identifier
        weights: {donor_id: weight} dict
        quality_score: fit quality
    """
    # Get pre-treatment outcomes for treatment and donors
    treatment_outcomes = get_pretreatment_outcomes(path_id)
    donor_outcomes = {d: get_pretreatment_outcomes(d) for d in donor_pool}

    # Solve for optimal weights (simplified - use scipy.optimize in production)
    weights = optimize_synthetic_weights(treatment_outcomes, donor_outcomes)

    # Compute match quality (R² of pre-treatment fit)
    synthetic_outcomes = np.sum(
        [donor_outcomes[d] * w for d, w in weights.items()], axis=0
    )
    quality_score = 1 - np.mean(
        (treatment_outcomes - synthetic_outcomes) ** 2
    ) / np.var(treatment_outcomes)

    cohort_id = f"control_synth_{hash(path_id) % 1000}"

    # Store synthetic control
    query = f"""
    INSERT INTO `{PROJECT_ID}.{DATASET_ID}.control_cohorts`
    (cohort_id, path_id, method, synthetic_weights, quality_score)
    VALUES (
        '{cohort_id}',
        '{path_id}',
        'synthetic_control',
        JSON '{weights}',
        {quality_score}
    )
    """
    client.query(query).result()

    logger.info(
        f"✓ Created synthetic control for {path_id}: {cohort_id} (quality={quality_score:.2f})"
    )
    return cohort_id, weights, quality_score


def create_matched_control(
    path_id: str, features: Dict[str, float], n_matches: int = 5
) -> Tuple[str, List[str], float]:
    """
    Propensity score matching: find similar untreated units.

    Returns:
        cohort_id: identifier
        matched_paths: list of matched path_ids
        quality_score: covariate balance
    """
    # Get propensity scores (probability of treatment given features)
    propensity_scores = compute_propensity_scores(features)

    # Find nearest neighbors in propensity score space
    matched_paths = find_nearest_neighbors(path_id, propensity_scores, n_matches)

    # Check covariate balance
    quality_score = check_covariate_balance(path_id, matched_paths, features)

    cohort_id = f"control_matched_{hash(path_id) % 1000}"

    # Store matched control
    for match_id in matched_paths:
        query = f"""
        INSERT INTO `{PROJECT_ID}.{DATASET_ID}.control_cohorts`
        (cohort_id, path_id, method, quality_score, props)
        VALUES (
            '{cohort_id}',
            '{match_id}',
            'propensity_match',
            {quality_score},
            JSON '{{"matched_to": "{path_id}"}}'
        )
        """
        client.query(query).result()

    logger.info(
        f"✓ Created matched control for {path_id}: {n_matches} matches (quality={quality_score:.2f})"
    )
    return cohort_id, matched_paths, quality_score


# ============================================================================
# Helper Functions
# ============================================================================


def get_available_regions() -> List[str]:
    """Fetch available geographic regions from data."""
    query = f"""
    SELECT DISTINCT 
        JSON_EXTRACT_SCALAR(features, '$.region') AS region
    FROM `{PROJECT_ID}.{DATASET_ID}.events_forecast`
    WHERE JSON_EXTRACT_SCALAR(features, '$.region') IS NOT NULL
    """
    result = client.query(query).result()
    return [row.region for row in result]


def compute_balance_score(treatment: List[str], control: List[str]) -> float:
    """
    Simplified balance check: compare treatment/control sizes.
    In production: use standardized mean difference (SMD).
    """
    if not treatment or not control:
        return 0.0

    # Placeholder: assume uniform region sizes
    ratio = len(control) / len(treatment)
    balance = 1.0 - abs(1.0 - ratio)
    return max(0.0, min(1.0, balance))


def get_pretreatment_outcomes(path_id: str, days_before: int = 14) -> np.ndarray:
    """Get outcome time series before treatment."""
    query = f"""
    SELECT value
    FROM `{PROJECT_ID}.{DATASET_ID}.events_outcome`
    WHERE path_id = '{path_id}'
      AND t2 < (SELECT MIN(t1) FROM `{PROJECT_ID}.{DATASET_ID}.events_action` WHERE path_id = '{path_id}')
      AND t2 >= TIMESTAMP_SUB(
        (SELECT MIN(t1) FROM `{PROJECT_ID}.{DATASET_ID}.events_action` WHERE path_id = '{path_id}'),
        INTERVAL {days_before} DAY
      )
    ORDER BY t2
    """
    result = client.query(query).result()
    return np.array([row.value for row in result])


def optimize_synthetic_weights(
    treatment: np.ndarray, donors: Dict[str, np.ndarray]
) -> Dict[str, float]:
    """
    Solve for synthetic control weights.
    Simplified version - use cvxpy or scipy.optimize in production.
    """
    from scipy.optimize import minimize

    donor_ids = list(donors.keys())
    donor_matrix = np.column_stack([donors[d] for d in donor_ids])

    # Objective: minimize ||treatment - donors @ weights||^2
    def objective(w):
        return np.sum((treatment - donor_matrix @ w) ** 2)

    # Constraints: weights sum to 1, non-negative
    constraints = [
        {"type": "eq", "fun": lambda w: np.sum(w) - 1},
    ]
    bounds = [(0, 1) for _ in donor_ids]

    result = minimize(
        objective,
        x0=np.ones(len(donor_ids)) / len(donor_ids),
        bounds=bounds,
        constraints=constraints,
    )

    return {donor_ids[i]: w for i, w in enumerate(result.x) if w > 0.01}


def compute_propensity_scores(features: Dict[str, float]) -> Dict[str, float]:
    """
    Estimate propensity scores using logistic regression.
    Placeholder - use sklearn in production.
    """
    # Simplified: random scores for demo
    query = f"""
    SELECT path_id
    FROM `{PROJECT_ID}.{DATASET_ID}.events_forecast`
    LIMIT 100
    """
    result = client.query(query).result()
    return {row.path_id: np.random.random() for row in result}


def find_nearest_neighbors(
    path_id: str, scores: Dict[str, float], n_matches: int
) -> List[str]:
    """Find N nearest neighbors in propensity score space."""
    target_score = scores.get(path_id, 0.5)
    distances = {
        pid: abs(score - target_score)
        for pid, score in scores.items()
        if pid != path_id
    }
    sorted_matches = sorted(distances.items(), key=lambda x: x[1])
    return [pid for pid, _ in sorted_matches[:n_matches]]


def check_covariate_balance(
    treatment_id: str, control_ids: List[str], features: Dict[str, float]
) -> float:
    """
    Check covariate balance using standardized mean difference.
    Quality score: 1.0 = perfect balance, 0.0 = no balance.
    """
    # Simplified: assume good balance if we found matches
    return 0.85 if control_ids else 0.0


if __name__ == "__main__":
    # Example usage
    path_id = "path_0001"

    # Method 1: Geo holdout
    cohort_id, quality = create_geo_holdout(
        path_id, treatment_regions=["US-CA", "US-NY"], test_ratio=0.3
    )
    print(f"Geo holdout: {cohort_id} (quality={quality:.2f})")

    # Method 2: Synthetic control (requires donor pool)
    # cohort_id, weights, quality = create_synthetic_control(...)

    # Method 3: Propensity matching
    # cohort_id, matches, quality = create_matched_control(...)
← All docsView source on GitHub →