Matplotlib - 基于数据的动态(条形)图表高度?

zwe*_*wer 6 python matplotlib

在与 matplotlib 斗争了比我愿意承认的时间更长的时间之后,我试图做一些在我用过的几乎所有其他绘图库中轻而易举的事情,我决定向 Stackiverse 寻求一些见解。简而言之,我需要的是创建多个水平条形图,所有条形图共享 x 轴,y 轴上具有不同数量的值,并且所有条形具有相同的高度,而图表本身将根据条形图的数量进行调整条目。我需要绘制的简化数据结构如下:

[
    {"name": "Category 1", "entries": [
        {"name": "Entry 1", "value": 5},
        {"name": "Entry 2", "value": 2},
    ]},
    {"name": "Category 2", "entries": [
        {"name": "Entry 1", "value": 1},
    ]},
    {"name": "Category 3", "entries": [
        {"name": "Entry 1", "value": 1},
        {"name": "Entry 2", "value": 10},
        {"name": "Entry 3", "value": 4},
    ]},    
]
Run Code Online (Sandbox Code Playgroud)

我最接近我想要的结果是使用:

import matplotlib.pyplot as plt

def plot_data(data):
    total_categories = len(data)  # holds how many charts to create
    max_values = 1  # holds the maximum number of bars to create
    for category in data:
        max_values = max(max_values, len(category["entries"]))
    fig = plt.figure(1)
    ax = None
    for index, category in enumerate(data):
        entries = []
        values = []
        for entry in category["entries"]:
            entries.append(entry["name"])
            values.append(entry["value"])
        if not entries:
            continue  # do not create empty charts
        y_ticks = range(1, len(entries) + 1)
        ax = fig.add_subplot(total_categories, 1, index + 1, sharex=ax)
        ax.barh(y_ticks, values)
        ax.set_ylim(0, max_values + 1)  # limit the y axis for fixed height
        ax.set_yticks(y_ticks)
        ax.set_yticklabels(entries)
        ax.invert_yaxis()
        ax.set_title(category["name"], loc="left")
    fig.tight_layout()
Run Code Online (Sandbox Code Playgroud)

无论有多少条目具有特定类别,这都将保持条形高度相同(至少在整个图中),这要归功于 y 限制 (set_ylim()) 设置为数据中条形的最高数量。然而,它也会在条目数量少于最大数量的类别中留下令人讨厌的空白。或者为了将所有内容放在视觉角度,我试图将其从实际变为预期:

图片链接

我尝试过通过 gridspec 和不同的比例来消除与条目数量相关的间隙,但最终看起来更加“倾斜”和不一致。我尝试生成多个图表并操纵图形大小,然后在后期处理中将它们缝合在一起,但我找不到一种方法可以可靠地使条形高度保持不变,无论数据如何。我确信有一种方法可以从 matplotlib 中的一些模糊对象中提取精确缩放所需的指标,但在这一点上,如果我尝试跟踪整个绘图过程,恐怕我会继续进行另一次徒劳的追逐。

此外,如果单个子图可以围绕数据折叠,我怎样才能使数字根据数据增长?例如,如果我要向上述数据添加第四个类别,而不是让图形的高度通过另一个图表“增长”,它实际上会缩小所有图表以适应默认图形大小的所有内容。现在,我想我理解了带有轴单位的 matplotlib 背后的逻辑以及所有这些,并且我知道我可以设置图形大小以增加整体高度,但我不知道如何在图表中保持一致,即如何无论数据如何,条形高度都完全相同?

我真的需要手动绘制所有内容才能得到我想要的吗?如果是这样,我可能会转储整个 matplotlib 包并从头开始创建我自己的 SVG。事后看来,考虑到我在这方面花费的时间,我可能一开始就应该这样做,但现在我太顽固了,无法放弃(或者我是可怕的沉没成本谬论的受害者) )。

有任何想法吗?

谢谢

Imp*_*est 5

我认为同时具有相等的条形宽度(垂直方向的宽度)和不同的子图大小的唯一方法实际上是手动定位图中的轴。

为此,您可以指定条形图应有多大(以英寸为单位),以及子图之间的间距(以该条形图宽度为单位)。根据这些数字以及要绘制的数据量,您可以计算出图形的总高度(以英寸为单位)。fig.add_axes然后根据数据量和前一个子图中的间距来定位每个子图(通过)。这样你就很好地填充了情节。添加一组新数据将使数字更大。

data = [
    {"name": "Category 1", "entries": [
        {"name": "Entry 1", "value": 5},
        {"name": "Entry 2", "value": 2},
    ]},
    {"name": "Category 2", "entries": [
        {"name": "Entry 1", "value": 1},
    ]},
    {"name": "Category 3", "entries": [
        {"name": "Entry 1", "value": 1},
        {"name": "Entry 2", "value": 10},
        {"name": "Entry 3", "value": 4},
    ]}, 
    {"name": "Category 4", "entries": [
        {"name": "Entry 1", "value": 6},
    ]},
]

import matplotlib.pyplot as plt
import numpy as np

def plot_data(data,
              barwidth = 0.2, # inch per bar
              spacing = 3,    # spacing between subplots in units of barwidth
              figx = 5,       # figure width in inch
              left = 4,       # left margin in units of bar width
              right=2):       # right margin in units of bar width

    tc = len(data) # "total_categories", holds how many charts to create
    max_values = []  # holds the maximum number of bars to create
    for category in data:
        max_values.append( len(category["entries"]))
    max_values = np.array(max_values)

    # total figure height:
    figy = ((np.sum(max_values)+tc) + (tc+1)*spacing)*barwidth #inch

    fig = plt.figure(figsize=(figx,figy))
    ax = None
    for index, category in enumerate(data):
        entries = []
        values = []
        for entry in category["entries"]:
            entries.append(entry["name"])
            values.append(entry["value"])
        if not entries:
            continue  # do not create empty charts
        y_ticks = range(1, len(entries) + 1)
        # coordinates of new axes [left, bottom, width, height]
        coord = [left*barwidth/figx, 
                 1-barwidth*((index+1)*spacing+np.sum(max_values[:index+1])+index+1)/figy,  
                 1-(left+right)*barwidth/figx,  
                 (max_values[index]+1)*barwidth/figy ] 

        ax = fig.add_axes(coord, sharex=ax)
        ax.barh(y_ticks, values)
        ax.set_ylim(0, max_values[index] + 1)  # limit the y axis for fixed height
        ax.set_yticks(y_ticks)
        ax.set_yticklabels(entries)
        ax.invert_yaxis()
        ax.set_title(category["name"], loc="left")


plot_data(data)
plt.savefig(__file__+".png")
plt.show()
Run Code Online (Sandbox Code Playgroud)

在此输入图像描述