import os
import numpy as np
import arviz as az
import pandas as pd
import matplotlib.pyplot as plt
import logging
import xarray
def extract_value(x):
"""Helper function to safely extract values from xarray/numpy objects recursively"""
logging.debug(f"extract_value: START - Input x: {x}, type: {type(x)}")
if isinstance(x, xarray.Dataset):
logging.error(f"extract_value received a Dataset: {x}. This is unexpected for a single diagnostic value. Returning None.")
logging.debug(f"extract_value: END - Returning None (input was Dataset)")
return None
# Try to get to a numpy array if it's xarray.DataArray
if isinstance(x, xarray.DataArray): # DataArray objects will have a .data attribute
logging.debug(f"extract_value: Input is DataArray. Accessing .data attribute.")
x_data = x.data
# If DataArray's .data attribute is itself a Dataset, this is an issue.
if isinstance(x_data, xarray.Dataset):
logging.error(f"extract_value: DataArray's .data attribute is a Dataset: {x_data}. Input DataArray was {x}. Returning None.")
logging.debug(f"extract_value: END - Returning None (DataArray.data was Dataset)")
return None
x = x_data # Proceed with x_data, which should be ndarray or scalar
logging.debug(f"extract_value: x is now x_data: {x}, type: {type(x)}")
# At this point, x should be a numpy array or already a Python/numpy scalar, not a Dataset or DataArray
if isinstance(x, np.ndarray):
logging.debug(f"extract_value: Input is ndarray. Size: {x.size}")
if x.size == 1:
try:
item_val = x.item() # Extract the single item from the 0-dim array
logging.debug(f"extract_value: ndarray.item() extracted: {item_val}, type: {type(item_val)}")
# Check if the extracted item is itself a complex xarray object
if isinstance(item_val, xarray.Dataset):
logging.error(f"extract_value: item extracted from numpy array is a Dataset: {item_val}. Original array was {x}. Returning None.")
logging.debug(f"extract_value: END - Returning None (item_val was Dataset)")
return None
if isinstance(item_val, xarray.DataArray):
# Recursively call extract_value if we get a DataArray from .item()
# This handles cases like np.array(xr.DataArray(value))
logging.warning(f"extract_value: item extracted from numpy array is a DataArray: {item_val}. Original array was {x}. Recursively calling extract_value.")
return extract_value(item_val) # Recursive call will have its own debug logs
logging.debug(f"extract_value: Attempting float(item_val) where item_val is: {item_val}, type: {type(item_val)}")
result = float(item_val) # Convert the (now confirmed scalar) item to Python float
logging.debug(f"extract_value: END - Returning float(item_val): {result}")
return result
except Exception as e_item:
item_val_type = type(item_val) if 'item_val' in locals() else 'unknown (item_val not defined)'
logging.error(f"Error calling .item() and/or float() on numpy array {x} (item_val type: {item_val_type}): {e_item}. Returning None.")
logging.debug(f"extract_value: END - Returning None (exception in ndarray.item()/float())")
return None
else:
logging.error(f"extract_value received a multi-element numpy array: {x}. Expected scalar. Returning None.")
logging.debug(f"extract_value: END - Returning None (multi-element ndarray)")
return None
# If x is already a Python scalar (int, float) or numpy scalar type
if isinstance(x, (int, float, np.number)):
logging.debug(f"extract_value: Input is scalar (int, float, np.number). Value: {x}, type: {type(x)}")
try:
logging.debug(f"extract_value: Attempting float(x) where x is: {x}, type: {type(x)}")
result = float(x) # Ensure it's a float
logging.debug(f"extract_value: END - Returning float(x): {result}")
return result
except Exception as e_float:
logging.error(f"Error converting {x} (type: {type(x)}) to float: {e_float}. Returning None.")
logging.debug(f"extract_value: END - Returning None (exception in scalar float())")
return None
logging.error(f"extract_value could not process input {x} (type: {type(x)}) to a scalar float. Returning None.")
logging.debug(f"extract_value: END - Returning None (unhandled type)")
return None
def save_elpd_to_csv(traces, output_dir):
elpd_data = []
for model_name, trace in traces.items():
try:
if hasattr(trace, 'log_likelihood'):
for var_name in trace.log_likelihood.data_vars:
try:
elpd = az.loo(trace, pointwise=True, var_name=var_name)
elpd_data.append({
'Model': model_name,
'Variable': var_name,
'ELPD': extract_value(elpd.elpd_loo),
'P_LOO': extract_value(elpd.p_loo)
})
except Exception as e:
logging.error(f"Error computing ELPD for {var_name} in {model_name}: {str(e)}")
else:
logging.warning(f"No log likelihood found in trace for model: {model_name}")
except Exception as e:
logging.error(f"Error processing model {model_name} for ELPD: {str(e)}")
if elpd_data:
elpd_df = pd.DataFrame(elpd_data)
elpd_df.to_csv(os.path.join(output_dir, 'elpd_comparison.csv'), index=False)
else:
logging.warning("No ELPD data to save.")
def save_diagnostics_to_csv(traces, output_dir):
for model_name, trace in traces.items():
try:
# Calculate summary only for parameters (exclude derived probabilities like p_...)
if hasattr(trace, 'posterior'):
param_vars = [v for v in trace.posterior.data_vars if not str(v).startswith('p_')]
if param_vars:
summary_df = az.summary(trace, var_names=param_vars, hdi_prob=0.95)
# The summary for scalar parameters should already be scalar, no need for extract_value loop
else:
logging.warning(f"No non-'p_' variables found in posterior for model {model_name}. Skipping summary CSV.")
continue # Skip to next model if no parameters found
else:
logging.warning(f"No posterior group found in trace for model {model_name}. Skipping summary CSV.")
continue # Skip to next model if no posterior
summary_df.to_csv(os.path.join(output_dir, f'{model_name}_mcmc_summary.csv'), index=True) # Renamed file for clarity
except Exception as e:
logging.error(f"Error computing or saving summary for {model_name}: {str(e)}")
def plot_trace(trace, model_name, output_dir):
try:
if hasattr(trace, 'posterior'):
param_vars = [v for v in trace.posterior.data_vars if not str(v).startswith('p_')]
if param_vars:
az.plot_trace(trace, var_names=param_vars)
plt.tight_layout()
plt.savefig(os.path.join(output_dir, f'{model_name}_trace_plot.png')) # Include model name
else:
logging.warning(f"No parameters found to plot trace for {model_name}.")
else:
logging.warning(f"No posterior group found to plot trace for {model_name}.")
except Exception as e:
logging.error(f"Error generating trace plot for {model_name}: {e}")
finally:
plt.close() # Correct indentation
def plot_posterior(trace, model_name, output_dir):
try:
if hasattr(trace, 'posterior'):
param_vars = [v for v in trace.posterior.data_vars if not str(v).startswith('p_')]
if param_vars:
az.plot_posterior(trace, var_names=param_vars)
plt.tight_layout()
plt.savefig(os.path.join(output_dir, f'{model_name}_posterior_plot.png')) # Include model name
else:
logging.warning(f"No parameters found to plot posterior for {model_name}.")
else:
logging.warning(f"No posterior group found to plot posterior for {model_name}.")
except Exception as e:
logging.error(f"Error generating posterior plot for {model_name}: {e}")
finally:
plt.close() # Correct indentation
def plot_energy(trace, model_name, output_dir):
try:
# Energy plot uses sample_stats, doesn't need var_names filtering usually
if hasattr(trace, 'sample_stats'):
az.plot_energy(trace)
plt.tight_layout()
plt.savefig(os.path.join(output_dir, f'{model_name}_energy_plot.png')) # Include model name
else:
logging.warning(f"No sample_stats group found to plot energy for {model_name}.")
except Exception as e:
logging.error(f"Error generating energy plot for {model_name}: {e}")
finally:
plt.close() # Correct indentation
[docs]
def run_diagnostics(idata):
"""
Run comprehensive MCMC diagnostics on an InferenceData object.
Parameters
----------
idata : arviz.InferenceData
The InferenceData object containing the MCMC samples
Returns
-------
dict
Dictionary containing various diagnostic metrics
"""
diagnostics = {
'rhat': {},
'ess': {},
'mcse': {},
'divergences': 0,
'warnings': []
}
try:
# Check if we have posterior samples
if not hasattr(idata, 'posterior'):
raise ValueError("No posterior samples found in InferenceData object")
# Calculate R-hat values for parameters only (exclude derived probabilities like p_...)
param_vars = [v for v in idata.posterior.data_vars if not str(v).startswith('p_')]
if not param_vars:
logging.warning("No non-'p_' variables found in posterior for diagnostics.")
return diagnostics # Return early if no parameters found
rhat_data = az.rhat(idata, var_names=param_vars)
# rhat_data is an xarray.Dataset. Iterate through its data_vars.
for var_name_str in rhat_data.data_vars:
try:
# Each data_var in the rhat_data Dataset should be a DataArray
actual_diagnostic_dataarray = rhat_data[var_name_str]
if not isinstance(actual_diagnostic_dataarray, xarray.DataArray):
logging.warning(f"R-hat for {var_name_str}: Expected DataArray, got {type(actual_diagnostic_dataarray)}. Skipping.")
continue
logging.debug(f"run_diagnostics (R-hat): Calling extract_value for {var_name_str} with actual_diagnostic_dataarray: {actual_diagnostic_dataarray}, type: {type(actual_diagnostic_dataarray)}")
extracted_val = extract_value(actual_diagnostic_dataarray)
logging.debug(f"run_diagnostics (R-hat): extract_value returned for {var_name_str}: {extracted_val}, type: {type(extracted_val)}")
if extracted_val is None or not isinstance(extracted_val, (int, float)):
logging.error(f"Could not extract scalar for R-hat of {var_name_str}. Value: {actual_diagnostic_dataarray}, Extracted: {extracted_val}")
continue
logging.debug(f"run_diagnostics (R-hat): Attempting float(extracted_val) for {var_name_str}, where extracted_val is: {extracted_val}, type: {type(extracted_val)}")
value = float(extracted_val)
logging.debug(f"run_diagnostics (R-hat): Converted value for {var_name_str}: {value}, type: {type(value)}")
diagnostics['rhat'][str(var_name_str)] = value
if value > 1.01:
diagnostics['warnings'].append(f"High R-hat ({value:.3f}) for {var_name_str}")
except Exception as e:
logging.error(f"Error processing R-hat for parameter {var_name_str}: {str(e)}")
# Calculate ESS for parameters only
ess_data = az.ess(idata, var_names=param_vars)
# ess_data is an xarray.Dataset. Iterate through its data_vars.
for var_name_str in ess_data.data_vars:
try:
actual_diagnostic_dataarray = ess_data[var_name_str]
if not isinstance(actual_diagnostic_dataarray, xarray.DataArray):
logging.warning(f"ESS for {var_name_str}: Expected DataArray, got {type(actual_diagnostic_dataarray)}. Skipping.")
continue
logging.debug(f"run_diagnostics (ESS): Calling extract_value for {var_name_str} with actual_diagnostic_dataarray: {actual_diagnostic_dataarray}, type: {type(actual_diagnostic_dataarray)}")
extracted_val = extract_value(actual_diagnostic_dataarray)
logging.debug(f"run_diagnostics (ESS): extract_value returned for {var_name_str}: {extracted_val}, type: {type(extracted_val)}")
if extracted_val is None or not isinstance(extracted_val, (int, float)):
logging.error(f"Could not extract scalar for ESS of {var_name_str}. Value: {actual_diagnostic_dataarray}, Extracted: {extracted_val}")
continue
logging.debug(f"run_diagnostics (ESS): Attempting float(extracted_val) for {var_name_str}, where extracted_val is: {extracted_val}, type: {type(extracted_val)}")
value = float(extracted_val)
logging.debug(f"run_diagnostics (ESS): Converted value for {var_name_str}: {value}, type: {type(value)}")
diagnostics['ess'][str(var_name_str)] = value
if value < 400: # Standard threshold for ESS
diagnostics['warnings'].append(f"Low ESS ({value:.1f}) for {var_name_str}")
except Exception as e:
logging.error(f"Error processing ESS for parameter {var_name_str}: {str(e)}")
# Calculate MCSE for parameters only
mcse_data = az.mcse(idata, var_names=param_vars)
# mcse_data is an xarray.Dataset. Iterate through its data_vars.
for var_name_str in mcse_data.data_vars:
try:
actual_diagnostic_dataarray = mcse_data[var_name_str]
if not isinstance(actual_diagnostic_dataarray, xarray.DataArray):
logging.warning(f"MCSE for {var_name_str}: Expected DataArray, got {type(actual_diagnostic_dataarray)}. Skipping.")
continue
logging.debug(f"run_diagnostics (MCSE): Calling extract_value for {var_name_str} with actual_diagnostic_dataarray: {actual_diagnostic_dataarray}, type: {type(actual_diagnostic_dataarray)}")
extracted_val = extract_value(actual_diagnostic_dataarray)
logging.debug(f"run_diagnostics (MCSE): extract_value returned for {var_name_str}: {extracted_val}, type: {type(extracted_val)}")
if extracted_val is None or not isinstance(extracted_val, (int, float)):
logging.error(f"Could not extract scalar for MCSE of {var_name_str}. Value: {actual_diagnostic_dataarray}, Extracted: {extracted_val}")
continue
logging.debug(f"run_diagnostics (MCSE): Attempting float(extracted_val) for {var_name_str}, where extracted_val is: {extracted_val}, type: {type(extracted_val)}")
value = float(extracted_val)
logging.debug(f"run_diagnostics (MCSE): Converted value for {var_name_str}: {value}, type: {type(value)}")
diagnostics['mcse'][str(var_name_str)] = value
except Exception as e:
logging.error(f"Error processing MCSE for parameter {var_name_str}: {str(e)}")
# Check divergences
if hasattr(idata, 'sample_stats') and hasattr(idata.sample_stats, 'diverging'):
try:
val_da = idata.sample_stats.diverging.sum() # This is an xarray.DataArray (0-dim)
logging.debug(f"run_diagnostics (Divergences): val_da from diverging.sum(): {val_da}, type: {type(val_da)}")
# Removed buggy while loop. extract_value can handle a 0-dim DataArray.
extracted_val = extract_value(val_da)
logging.debug(f"run_diagnostics (Divergences): extract_value returned: {extracted_val}, type: {type(extracted_val)}")
if extracted_val is None or not isinstance(extracted_val, (int, float)): # Check if None or not number
logging.error(f"Could not extract scalar for divergences. Value: {val_da}, Extracted: {extracted_val}")
n_divergent = 0 # Default if extraction fails
else:
logging.debug(f"run_diagnostics (Divergences): Attempting int(extracted_val), where extracted_val is: {extracted_val}, type: {type(extracted_val)}")
n_divergent = int(extracted_val)
logging.debug(f"run_diagnostics (Divergences): n_divergent: {n_divergent}, type: {type(n_divergent)}")
diagnostics['divergences'] = n_divergent
if n_divergent > 0:
diagnostics['warnings'].append(f"Found {n_divergent} divergent transitions")
except Exception as e:
logging.error(f"Error checking divergences: {str(e)}")
# Add sampling statistics
try:
diagnostics['n_chains'] = int(idata.posterior.sizes['chain'])
diagnostics['n_draws'] = int(idata.posterior.sizes['draw'])
if diagnostics['n_chains'] < 4:
diagnostics['warnings'].append(f"Only {diagnostics['n_chains']} chains used. Consider using at least 4 chains.")
except Exception as e:
logging.error(f"Error computing sampling statistics: {str(e)}")
logging.info("MCMC diagnostics computed successfully")
for warning in diagnostics['warnings']:
logging.warning(warning)
except Exception as e:
logging.error(f"Error computing MCMC diagnostics: {str(e)}")
diagnostics['warnings'].append(f"Failed to compute some diagnostics: {str(e)}")
return diagnostics