代码之家  ›  专栏  ›  技术社区  ›  lo tolmencre

对齐顶轴和底轴的勾号坐标

  •  2
  • lo tolmencre  · 技术社区  · 7 年前

    我想构建一个图形,在其底部x轴上有一组记号,在其顶部x轴上有另一组与底部记号对齐的记号。特别是在我的例子中,这些是批次和时代。每 n

    import numpy as np
    import matplotlib.pyplot as plt
    
    batches = np.arange(1,101)
    epoch_ends = batches[[(i*10)-1 for i in range(1,11)]]
    accuracy = np.apply_along_axis(arr=batches, axis=0, func1d=lambda x : x/len(batches))
    loss = np.apply_along_axis(arr=batches, axis=0, func1d=lambda x : 1 - (x/len(batches)))
    
    fig, ax1 = plt.subplots( nrows=1, ncols=1 )
    ax2 = ax1.twinx()
    ax3 = ax1.twiny()
    
    ax1.set_xlabel('batches')
    ax1.set_xticks(np.arange(1, len(batches)+1, 9))
    ax1.set_ylabel('accuracy')
    ax1.grid()
    
    ax2.set_ylabel('loss')
    ax2.set_yticklabels(np.linspace(3, 10, 9))
    
    ax3.set_xlabel('epochs')
    ax3.set_xticks(epoch_ends)
    ax3.set_xticklabels(range(1, len(epoch_ends)+1))
    
    acc_plt = ax1.plot(batches, accuracy, color='red', label='acc')
    loss_plt = ax2.plot(batches, loss, color='blue', label='loss')
    
    lns = acc_plt+loss_plt
    labs = [l.get_label() for l in lns]
    ax1.legend(lns, labs, loc=2)
    
    plt.show()
    

    batches epoch_ends 分别是这样的

    [  1   2   3   4   5   6   7   8   9  10  11  12  13  14  15  16  17  18
      19  20  21  22  23  24  25  26  27  28  29  30  31  32  33  34  35  36
      37  38  39  40  41  42  43  44  45  46  47  48  49  50  51  52  53  54
      55  56  57  58  59  60  61  62  63  64  65  66  67  68  69  70  71  72
      73  74  75  76  77  78  79  80  81  82  83  84  85  86  87  88  89  90
      91  92  93  94  95  96  97  98  99 100]
    [ 10  20  30  40  50  60  70  80  90 100]
    

    因此,我希望历元记号1与批次x-coorrdiante 10、2与20对齐,以此类推。

    enter image description here

    我需要在代码中更改什么才能使其正常工作?

    2 回复  |  直到 7 年前
        1
  •  2
  •   Sheldore    7 年前

    这里有一种对齐它们的方法。想法如下:

    • 首先使用以下工具在下x轴上绘制数据:( ax1 )
    • 然后使用将上x轴的限制设置为与下x轴相同 ax3.set_xlim(ax1.get_xlim())
    • 然后在对应于下x轴值(10、20、30、…、90、100)的位置设置上x轴的刻度
    • 最后,使用重命名记号标签 ax3.set_xticklabels()

    代码如下:我正在用注释替换代码中已经存在的部分 # .

    # imports and define data and compute accuracy and loss here
    
    # Initiate figure and axis objects here
    
    ax1.set_xlabel('batches')
    ax1.set_xticks(np.arange(1, len(batches)+1, 9))
    ax1.set_ylabel('accuracy')
    ax1.grid()
    
    acc_plt = ax1.plot(batches, accuracy, color='red', label='acc')
    loss_plt = ax2.plot(batches, loss, color='blue', label='loss')
    
    ax2.set_ylabel('loss')
    ax2.set_yticklabels(np.linspace(3, 10, 9))
    
    new_tick_locations = np.arange(1, 11)*10
    new_tick_labels = range(1, 11)
    
    ax3.set_xlabel('epochs')
    ax3.set_xlim(ax1.get_xlim())
    ax3.set_xticks(new_tick_locations)
    ax3.set_xticklabels(new_tick_labels)
    
    # Set legends and show the plot
    

    enter image description here

        2
  •  1
  •   Owen    7 年前

    (新行已注释)

    import numpy as np
    import matplotlib.pyplot as plt
    
    batches = np.arange(1,101)
    epoch_ends = batches[[(i*10)-1 for i in range(1,11)]]
    accuracy = np.apply_along_axis(arr=batches, axis=0, func1d=lambda x : x/len(batches))
    loss = np.apply_along_axis(arr=batches, axis=0, func1d=lambda x : 1 - (x/len(batches)))
    
    fig, ax1 = plt.subplots( nrows=1, ncols=1 )
    ax2 = ax1.twinx()
    ax3 = ax1.twiny()
    
    ax1.set_xlabel('batches')
    #ax1.set_xticks(np.arange(1, len(batches)+1, 9))
    ax1.set_xlim(0, 100)                               # Set xlim
    ax1.set_ylabel('accuracy')
    ax1.grid()
    
    ax2.set_ylabel('loss')
    ax2.set_yticklabels(np.linspace(3, 10, 9))
    
    ax3.set_xlabel('epochs')
    #ax3.set_xticks(epoch_ends)
    #ax3.set_xticklabels(range(1, len(epoch_ends)+1))
    ax3.set_xlim(0, 10)                                # Set xlim
    
    acc_plt = ax1.plot(batches, accuracy, color='red', label='acc')
    loss_plt = ax2.plot(batches, loss, color='blue', label='loss')
    
    lns = acc_plt+loss_plt
    labs = [l.get_label() for l in lns]
    ax1.legend(lns, labs, loc=2)
    
    plt.show()