You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用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 line x1 + x2 = 1 where x1, 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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.15 04:29:16