使用Matplotlib绘制Dirichlet分布热力图出现异常图像的原因
Hey there! Let's break down why your heatmap isn't turning out as expected, and fix it up.
First, let's look at what your code is actually doing:
np.random.dirichlet(alpha=[0.3, 0.7], size=1000)generates a (1000, 2) array. Every row in this array sums to 1 (that's the defining property of the Dirichlet distribution on a 2-simplex—think of it as points lying on the linex1 + x2 = 1wherex1, x2 ≥ 0).plt.imshow()is designed to visualize 2D matrices (images). When you pass a (1000,2) array to it, it treats each row as a horizontal line of 2 pixels, resulting in a weird "two vertical strips" plot. This has nothing to do with the actual shape of the Dirichlet distribution—you're using the wrong tool for the job here.
Fix 1: Visualize 2D Dirichlet with a Histogram
Since a 2D Dirichlet reduces to a 1D distribution (because x2 = 1 - x1), the simplest way to show its shape is with a histogram of one of the variables:
import numpy as np import matplotlib.pyplot as plt a = np.random.dirichlet(alpha=[0.3, 0.7], size=1000) # Grab the first variable (the second would just be its mirror) x = a[:, 0] plt.hist(x, bins=30, density=True, alpha=0.6, color='#e74c3c') plt.xlabel('$x_1$') plt.ylabel('Density') plt.title('Dirichlet Distribution (α = [0.3, 0.7])') plt.show()
Fix 2: Plot Samples on the 2-Simplex
If you want to see where the samples lie geometrically, you can plot them on the line x1 + x2 = 1:
import numpy as np import matplotlib.pyplot as plt a = np.random.dirichlet(alpha=[0.3, 0.7], size=1000) x1 = a[:, 0] x2 = a[:, 1] plt.scatter(x1, x2, s=6, alpha=0.5, color='#3498db') # Draw the simplex boundary plt.plot([0, 1], [1, 0], 'k--', linewidth=1) plt.xlim(0, 1) plt.ylim(0, 1) plt.xlabel('$x_1$') plt.ylabel('$x_2$') plt.title('Dirichlet Samples on the 2-Simplex') plt.show()
If You Really Want a Heatmap (Use 3D Dirichlet)
Heatmaps make sense for higher-dimensional Dirichlet distributions (3D or more), where you can plot the joint distribution of the first two variables (since the third is 1 - x1 - x2). Here's how to do that:
import numpy as np import matplotlib.pyplot as plt # Generate samples from a 3D Dirichlet distribution a = np.random.dirichlet(alpha=[0.3, 0.7, 1.0], size=10000) x1 = a[:, 0] x2 = a[:, 1] # Create a 2D histogram heatmap plt.hist2d(x1, x2, bins=35, cmap='hot') plt.colorbar(label='Sample Count') plt.xlabel('$x_1$') plt.ylabel('$x_2$') plt.title('Heatmap of 3D Dirichlet Distribution') plt.show()
To recap:
- A 2D Dirichlet is a 1D distribution, so
imshow(a 2D image tool) is not the right choice here. - Use histograms or simplex scatter plots for 2D Dirichlet.
- Reserve heatmaps for 3D+ Dirichlet distributions, where you can visualize the joint density of two variables.
内容的提问来源于stack exchange,提问作者user8270077
相关产品推荐
相关产品推荐

