import numpy as np
import matplotlib.pyplot as plt
np.random.seed(42)
# Create two highly correlated features (ρ = 0.9)
p = 0.9
mean = [0, 0]
cov = [[1, p], [p, 1]]
n = 100
x1, x2 = np.random.multivariate_normal(mean, cov, n).T
# Point to explain
point = (-1.7, -1.7)
m = 15 # number of samples
# Marginal sampling: sample x2 independently (ignores correlation)
x2_marg = np.random.choice(x2, size=m)
# Conditional sampling: sample x2 given x1 (respects correlation)
x2_cond = np.random.normal(loc=p*point[0], scale=np.sqrt(1-p**2), size=m)
# Create side-by-side plots
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(14, 6))
# Left: Marginal sampling (SHAP default)
ax1.scatter(x1, x2, color='black', alpha=0.2, s=50, label='Training data')
ax1.scatter(np.repeat(point[0], m), x2_marg, color='blue', s=100, alpha=0.7, label='Marginal samples')
ax1.scatter(point[0], point[1], color='red', s=200, marker='*', label='Point to explain', zorder=10)
ax1.set_xlabel('Feature X₁', fontsize=14)
ax1.set_ylabel('Feature X₂', fontsize=14)
ax1.set_title('Marginal Sampling (SHAP Default)\n⚠️ Creates unrealistic data points', fontsize=16, fontweight='bold')
ax1.legend(fontsize=12)
ax1.grid(alpha=0.3)
ax1.set_xlim(-2.5, 2.5)
ax1.set_ylim(-2.5, 2.5)
# Right: Conditional sampling
ax2.scatter(x1, x2, color='black', alpha=0.2, s=50, label='Training data')
ax2.scatter(np.repeat(point[0], m), x2_cond, color='green', s=100, alpha=0.7, label='Conditional samples')
ax2.scatter(point[0], point[1], color='red', s=200, marker='*', label='Point to explain', zorder=10)
ax2.set_xlabel('Feature X₁', fontsize=14)
ax2.set_ylabel('Feature X₂', fontsize=14)
ax2.set_title('Conditional Sampling P(X₂|X₁)\n✓ Respects feature correlation', fontsize=16, fontweight='bold')
ax2.legend(fontsize=12)
ax2.grid(alpha=0.3)
ax2.set_xlim(-2.5, 2.5)
ax2.set_ylim(-2.5, 2.5)
plt.tight_layout()
plt.show()