Source code for src.utils.diagnostics

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