我有一个多标签数据集,我正在使用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)