Source code for src.utils.analysis

import pandas as pd
import numpy as np
import requests
import json
import os
from datetime import datetime
import logging

from src.utils.saving import ensure_output_dir, validate_inference_data, save_mcmc_diagnostics, save_parameter_summary, save_convergence_diagnostics, save_parameter_traces, save_parameter_comparison
from src.utils.visualization import (
    plot_forest_comparison, 
    plot_dag,
    plot_posterior_effects,
    plot_posterior_predictive,
    plot_conversion_funnel,
    plot_path_analysis,
    plot_probability_flow
)
from src.utils.diagnostics import run_diagnostics

logger = logging.getLogger(__name__)


[docs] def analyze_results_gemini(analysis_data, gemini_api_key, output_dir): """ Analyzes the given results using the Google Gemini API and returns a text analysis along with recommendations. Displays formatted results in terminal and saves to file. Parameters: analysis_data (dict): The results from the model. gemini_api_key (str): Your Google Gemini API key. output_dir (str): Directory to save the analysis text file. Returns: str: The analysis and recommendations from Gemini. """ # Create context-aware prompt prompt = f""" As an expert in Bayesian networks and customer journey analysis, analyze the following ConversionFlow results: Context: - ConversionFlow analyzes customer journeys through Bayesian networks - The model tracks key conversion events - The causal graph includes stages from session start through to final conversion Model Results: ``` {json.dumps(analysis_data, indent=2)} ``` """ # Add MCMC diagnostics to prompt mcmc_diag_path = os.path.join(output_dir, 'initial_model_mcmc_diagnostics.csv') if os.path.exists(mcmc_diag_path): with open(mcmc_diag_path, 'r') as f: mcmc_diagnostics = f.read() prompt += f""" MCMC Diagnostics: ``` {mcmc_diagnostics} ``` """ prompt += f""" Please provide a comprehensive analysis, including: 1. Model Convergence Analysis: - MCMC diagnostics interpretation (R-hat, effective sample size) - Parameter convergence assessment - Any convergence issues or warnings 2. Customer Journey Insights: - Strongest causal relationships identified - Key conversion bottlenecks - Most influential touchpoints 3. Business Recommendations: - Specific areas for conversion optimization - Prioritized improvement opportunities 4. Model Improvement Suggestions: - Parameter adjustments - Prior distribution recommendations - Data quality considerations Format the analysis in clear sections with bullet points for actionability. """ print("\nPROCESSING RESULTS...\n") payload = { "contents": [{ "parts": [{"text": prompt}] }] } url = f"https://generativelanguage.googleapis.com/v1beta/models/gemini-1.5-flash:generateContent?key={gemini_api_key}" headers = {'Content-Type': 'application/json'} try: response = requests.post(url, headers=headers, data=json.dumps(payload)) response.raise_for_status() # Raise exception for bad status codes data = response.json() if not data.get('candidates'): raise requests.exceptions.RequestException("No candidates in API response") candidate = data['candidates'][0] if not candidate.get('content') or not candidate['content'].get('parts'): raise requests.exceptions.RequestException("Invalid response structure") analysis = candidate['content']['parts'][0].get('text', '') if not analysis: raise requests.exceptions.RequestException("No text content in API response") timestamp = datetime.now().strftime('%Y-%m-%d %H:%M:%S') # Format and display the analysis print("\n" + "="*80) print("CONVERSIONFLOW ANALYSIS RESULTS") print(f"Generated at: {timestamp}") print("="*80 + "\n") # Clean up and format the analysis clean_analysis = (analysis .replace('## ', '') # Remove markdown headers .replace('**', '') # Remove bold markers .replace('`', '') # Remove code markers .replace('\u2192', '->') # Replace unicode arrow .replace('\u03b2', 'beta') # Replace unicode beta ) # Split into sections and format sections = clean_analysis.split('\n\n') formatted_analysis = [] for section in sections: if section.strip(): # Add section header if section.startswith('Analysis of'): formatted_analysis.append(section) elif section.startswith('1.') or section.startswith('2.') or section.startswith('3.') or section.startswith('4.'): formatted_analysis.append("\n" + "-"*80) formatted_analysis.append(section) else: formatted_analysis.append(section) formatted_text = '\n'.join(formatted_analysis) # Print formatted analysis print(formatted_text) print("\n" + "="*80) # Save complete analysis to text file analysis_path = os.path.join(output_dir, 'model_analysis.txt') with open(analysis_path, 'w') as f: f.write("CONVERSIONFLOW ANALYSIS RESULTS\n") f.write(f"Generated at: {timestamp}\n") f.write("="*80 + "\n\n") f.write(formatted_text) f.write("\n" + "="*80 + "\n") logging.info(f"Analysis saved to: {analysis_path}") return analysis except requests.exceptions.RequestException as e: error_msg = f"Error: {str(e)}" print("\nCONVERSIONFLOW ANALYSIS ERROR") print("-"*80) print("-"*80 + "\n") return error_msg
[docs] def analyze_results(results, output_dir, config, preprocessing_stats, data): """Analyze and save results.""" try: # Ensure output directory exists ensure_output_dir(output_dir) logging.info(f"Ensured output directory exists: {output_dir}") analysis_data = { 'models': {}, 'preprocessing_stats': preprocessing_stats, 'metadata': { 'timestamp': datetime.now().isoformat(), 'config': { 'inference': config.get('inference', {}), 'model': config.get('model', {}) } } } for model_name, model_data in results.items(): if 'trace' not in model_data: logging.warning(f"No trace data found for {model_name}") continue idata = model_data['trace'] if not hasattr(idata, 'posterior'): logging.warning(f"No posterior samples found in trace for {model_name}") continue model_info = {'name': model_name} # Save all diagnostics and summaries try: # Validate inference data validate_inference_data(idata, "analyze_results") # Save diagnostics mcmc_diag_path = save_mcmc_diagnostics(idata, model_name, output_dir) param_summary_path = save_parameter_summary(idata, model_name, output_dir) convergence_path = save_convergence_diagnostics(idata, model_name, output_dir) # Save parameter traces and comparison trace_path = save_parameter_traces(idata, model_name, output_dir) if hasattr(idata, 'log_likelihood'): comparison_path = save_parameter_comparison(idata, model_name, output_dir) mcmc_summary_path = os.path.join(output_dir, f"{model_name}_mcmc_summary.json") warnings_path = os.path.join(output_dir, f"{model_name}_mcmc_warnings.txt") logging.info("Saved diagnostics files:") logging.info(f"- MCMC diagnostics: {mcmc_diag_path}") logging.info(f"- Parameter summary: {param_summary_path}") logging.info(f"- Convergence diagnostics: {convergence_path}") logging.info(f"- Parameter traces: {trace_path}") if hasattr(idata, 'log_likelihood'): logging.info(f"- Parameter comparison: {comparison_path}") except Exception as e: logging.error(f"Error saving diagnostics files: {str(e)}") continue # Initialize model_info substructures model_info['posterior_summary'] = None model_info['mcmc_diagnostics'] = {} # Ensure this dict exists model_info['convergence_metrics'] = None model_info['mcmc_summary'] = None model_info['mcmc_warnings'] = [] # Load parameter summary CSV for detailed stats if param_summary_path and os.path.exists(param_summary_path): try: param_df = pd.read_csv(param_summary_path, index_col=0) if not param_df.empty: r_hat_values = param_df['r_hat'].dropna().values ess_bulk_values = param_df['ess_bulk'].dropna().values ess_tail_values = param_df['ess_tail'].dropna().values if not (len(r_hat_values) == 0 or len(ess_bulk_values) == 0 or len(ess_tail_values) == 0): posterior_summary = {} for param_idx in param_df.index: posterior_summary[param_idx] = { 'mean': float(param_df.loc[param_idx, 'mean']), 'std': float(param_df.loc[param_idx, 'sd']), 'hdi_2.5%': float(param_df.loc[param_idx, 'hdi_2.5%']), 'hdi_97.5%': float(param_df.loc[param_idx, 'hdi_97.5%']), 'median': float(param_df.loc[param_idx, 'median']), 'diagnostics': { 'r_hat': float(param_df.loc[param_idx, 'r_hat']), 'ess_bulk': float(param_df.loc[param_idx, 'ess_bulk']), 'ess_tail': float(param_df.loc[param_idx, 'ess_tail']), 'mcse_mean': float(param_df.loc[param_idx, 'mcse_mean']), 'mcse_sd': float(param_df.loc[param_idx, 'mcse_sd']) } } model_info['posterior_summary'] = posterior_summary model_info['mcmc_diagnostics']['summary_stats'] = { 'r_hat_max': float(np.max(r_hat_values)) if r_hat_values.size > 0 else None, 'r_hat_min': float(np.min(r_hat_values)) if r_hat_values.size > 0 else None, 'ess_bulk_min': float(np.min(ess_bulk_values)) if ess_bulk_values.size > 0 else None, 'ess_bulk_mean': float(np.mean(ess_bulk_values)) if ess_bulk_values.size > 0 else None, 'ess_tail_min': float(np.min(ess_tail_values)) if ess_tail_values.size > 0 else None, 'ess_tail_mean': float(np.mean(ess_tail_values)) if ess_tail_values.size > 0 else None } model_info['convergence_metrics'] = { 'r_hat_max': float(np.max(r_hat_values)) if r_hat_values.size > 0 else None, 'ess_bulk_min': float(np.min(ess_bulk_values)) if ess_bulk_values.size > 0 else None, 'ess_tail_min': float(np.min(ess_tail_values)) if ess_tail_values.size > 0 else None, 'has_convergence_issues': any([ (np.max(r_hat_values) > 1.01 if r_hat_values.size > 0 else False), (np.min(ess_bulk_values) < 400 if ess_bulk_values.size > 0 else False), (np.min(ess_tail_values) < 400 if ess_tail_values.size > 0 else False) ]), 'problematic_parameters': [ param_idx for param_idx in param_df.index if (pd.notna(param_df.loc[param_idx, 'r_hat']) and param_df.loc[param_idx, 'r_hat'] > 1.01) or (pd.notna(param_df.loc[param_idx, 'ess_bulk']) and param_df.loc[param_idx, 'ess_bulk'] < 400) or (pd.notna(param_df.loc[param_idx, 'ess_tail']) and param_df.loc[param_idx, 'ess_tail'] < 400) ] } else: logging.warning(f"Empty MCMC diagnostic values in {param_summary_path} for {model_name}.") else: logging.warning(f"Parameter summary DataFrame is empty for {model_name} from {param_summary_path}.") except Exception as e: logging.error(f"Error processing parameter summary CSV for {model_name}: {str(e)}") # Load MCMC summary JSON (for overall stats like n_chains, divergences) if mcmc_summary_path and os.path.exists(mcmc_summary_path): try: with open(mcmc_summary_path, 'r') as f: mcmc_json_data = json.load(f) model_info['mcmc_summary'] = mcmc_json_data # Populate chain_stats from the JSON data if 'n_chains' in mcmc_json_data and 'n_draws_per_chain' in mcmc_json_data and 'total_samples' in mcmc_json_data: model_info['mcmc_diagnostics']['chain_stats'] = { 'n_chains': int(mcmc_json_data['n_chains']), 'n_draws': int(mcmc_json_data['n_draws_per_chain']), # Assuming n_draws means per chain from JSON 'total_samples': int(mcmc_json_data['total_samples']) } else: logging.warning(f"Chain stats not fully available in MCMC summary JSON for {model_name}") except Exception as e: logging.error(f"Error loading MCMC summary JSON for {model_name}: {str(e)}") # Load MCMC warnings text file if warnings_path and os.path.exists(warnings_path): try: with open(warnings_path, 'r') as f: model_info['mcmc_warnings'] = f.read().splitlines() except Exception as e: logging.error(f"Error loading MCMC warnings for {model_name}: {str(e)}") # Run comprehensive MCMC diagnostics and save additional analysis files try: # Run diagnostics diagnostics = run_diagnostics(idata) logging.info(f"Type of idata: {type(idata)}") logging.info(f"Data variables in idata.posterior: {list(idata.posterior.data_vars)}") logging.info(f"Coordination variables in idata.posterior: {list(idata.posterior.coords)}") logging.info(f"Sizes of idata.posterior: {idata.posterior.sizes}") # Save parameter traces save_parameter_traces(idata, model_name, output_dir) # Save parameter comparison if log likelihood is available if hasattr(idata, 'log_likelihood'): save_parameter_comparison(idata, model_name, output_dir) if model_info['mcmc_diagnostics'] is None: model_info['mcmc_diagnostics'] = {} model_info['mcmc_diagnostics'].update({ 'comprehensive_diagnostics': diagnostics }) # Add sampling stats from diagnostics model_info['sampling_stats'] = { 'num_chains': diagnostics.get('n_chains'), 'num_draws': diagnostics.get('n_draws'), 'num_divergences': diagnostics.get('divergences'), 'warnings': diagnostics.get('warnings', []) } except Exception as e: logging.error(f"Error computing comprehensive diagnostics for {model_name}: {str(e)}") model_info['sampling_stats'] = { 'num_chains': None, 'num_draws': None, 'num_divergences': None, 'warnings': [f"Failed to compute diagnostics: {str(e)}"] } analysis_data['models'][model_name] = model_info # Generate visualizations for this model try: # Ensure 'outcome' exists, otherwise default to an empty list or handle error outcome_nodes_config = config.get('model', {}).get('nodes', {}).get('outcome', []) if not isinstance(outcome_nodes_config, list): # Ensure it's a list logger.warning(f"Expected 'outcome' nodes to be a list, got {type(outcome_nodes_config)}. Defaulting to empty list for target_variables.") target_variables = [] else: target_variables = outcome_nodes_config # Generate all visualizations plot_dag(model_data.get('model', None), model_name, output_dir) plot_posterior_effects(idata, model_name, output_dir) # Get causal chain from config causal_chain = [(source, target) for source, target in config['model']['edges']] # Generate visualizations that need causal chain plot_conversion_funnel(data, causal_chain, output_dir) plot_path_analysis(idata, causal_chain, output_dir) plot_probability_flow(idata, causal_chain, output_dir) # Generate posterior predictive plots for each target variable # Use the already fetched target_variables list for target in target_variables: plot_posterior_predictive(idata, data, target, model_name, output_dir) # The recursive call analyze_model_results was here and has been removed. # Visualizations are generated, then the loop continues or exits. logging.info(f"Generated all visualizations for {model_name}") except Exception as e: logging.error(f"Error generating visualizations for {model_name}: {str(e)}") # Generate forest plot comparison now that all models are processed if results: plot_forest_comparison(results, output_dir) logging.info("Forest plot comparison saved") # Debug: Log the analysis data structure logging.info("Analysis data structure:") logging.info(json.dumps(analysis_data, indent=2)) # Save the raw analysis data for debugging with open(os.path.join(output_dir, 'raw_analysis_data.json'), 'w') as f: f.write(json.dumps(analysis_data, indent=2)) return analysis_data except Exception as e: logging.error(f"Error in analysis: {str(e)}") raise
# Alias for legacy compatibility analyze_model_results = analyze_results