【问题标题】:How to change the threshold of a prediction of multi-label classification using FASTAI library如何使用 FASTAI 库更改多标签分类预测的阈值
【发布时间】:2022-06-14 20:03:33
【问题描述】:

我有一个多标签数据集,我用它来训练我的模型,使用 Python 的 fast-ai 库,使用准确度函数作为指标,例如:

def accuracy_multi1(inp, targ, thresh=0.5, sigmoid=True):
    "Compute accuracy when 'inp' and 'targ' are the same size"
    if sigmoid: inp=inp.sigmoid()
    return ((inp>thresh) == targ.bool()).float().mean()

我的学习者是这样的:

learn = cnn_learner(dls, resnet50, metrics=partial(accuracy_multi1,thresh=0.1))
learn.fine_tune(2,base_lr=3e-2,freeze_epochs=2)

在训练我的模型之后,我想考虑使用参数的阈值来预测图像,但是方法learn.predict('img.jpg') 只考虑默认的thres=0.5。在下面的例子中,我的预测应该返回 True 代表“红色”、“衬衫”和“鞋子”,因为它们的概率高于 0.1(但鞋子低于 0.5,因此不被视为 True):

def printclasses(prediction,classes):
    print('Prediction:',prediction[0])
    for i in range(len(classes)):
        print(classes[i],':',bool(prediction[1][i]),'|',float(prediction[2][i]))

printclasses(learn.predict('rose.jpg'),dls.vocab)

输出:

Prediction: ['red', 'shirt']
black : False | 0.007274294272065163
blue : False | 0.0019288889598101377
brown : False | 0.005750810727477074
dress : False | 0.0028723080176860094
green : False | 0.005523672327399254
hoodie : False | 0.1325301229953766
pants : False | 0.009496113285422325
pink : False | 0.0037188702262938023
red : True | 0.9839697480201721
shirt : True | 0.5762518644332886
shoes : False | 0.2752271890640259
shorts : False | 0.0020902694668620825
silver : False | 0.0009014935349114239
skirt : False | 0.0030087409541010857
suit : False | 0.0006510693347081542
white : False | 0.001247694599442184
yellow : False | 0.0015280473744496703

当我对我引用的图像进行预测时,有没有办法施加阈值?看起来像这样的东西:

learn.predict('img.jpg',thresh=0.1)

【问题讨论】:

    标签: python deep-learning pytorch fast-ai


    【解决方案1】:

    我也遇到了同样的问题。我仍然对更好的解决方案感兴趣,但由于 accuracy_mult 似乎只在训练过程中提供对模型的用户友好评估(并且不参与预测),因此我为我的数据创建了一个解决方法。

    基本思想是将张量与实际预测(这是predict() 函数返回的三元组中的第三个条目)相结合,应用阈值并从词汇中获取相应的标签。

    def predict_labels(x, model, thresh=0.5):
      '''
      function to predict multi-labels in text (x)
    
      arguments:
      ----------
      x: the text to predict
      model: the trained learner
      thresh: thresh to indicate which labels should be included, fastai default is 0.5
    
      return:
      -------
      (str) predictions separated by blankspace
      '''
    
      # getting categories according to threshold
      preds = model.predict(x)[2] > thresh
      labels = model.dls.multi_categorize.vocab[preds]
    
      return ' '.join(labels) 
    
    

    【讨论】:

      猜你喜欢
      • 2019-09-02
      • 1970-01-01
      • 1970-01-01
      • 2021-04-21
      • 2018-11-13
      • 2015-11-14
      • 2021-07-19
      • 2022-06-28
      • 2014-03-12
      相关资源
      最近更新 更多