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