【发布时间】:2019-08-05 17:16:54
【问题描述】:
我一直在关注 spaCy 文本分类快速入门指南。
假设我有一个非常简单的数据集。
TRAIN_DATA = [
("beef", {"cats": {"POSITIVE": 1.0, "NEGATIVE": 0.0}}),
("apple", {"cats": {"POSITIVE": 0, "NEGATIVE": 1}})
]
我正在训练一个管道来对文本进行分类。它可以训练并且丢失率很低。
textcat = nlp.create_pipe("pytt_textcat", config={"exclusive_classes": True})
for label in ("POSITIVE", "NEGATIVE"):
textcat.add_label(label)
nlp.add_pipe(textcat)
optimizer = nlp.resume_training()
for i in range(10):
random.shuffle(TRAIN_DATA)
losses = {}
for batch in minibatch(TRAIN_DATA, size=8):
texts, cats = zip(*batch)
nlp.update(texts, cats, sgd=optimizer, losses=losses)
print(i, losses)
现在,我如何预测一个新的文本字符串是“正”还是“负”?
这将起作用:
doc = nlp(u'Pork')
print(doc.cats)
它为我们训练预测的每个类别提供一个分数。
但这似乎与文档不一致。 It says I should use a predict method 在原来的子类管道组件上。
但那行不通。
尝试textcat.predict('text') 或textcat.predict(['text']) 等。抛出:
AttributeError Traceback (most recent call last)
<ipython-input-29-39e0c6e34fd8> in <module>
----> 1 textcat.predict(['text'])
pipes.pyx in spacy.pipeline.pipes.TextCategorizer.predict()
AttributeError: 'str' object has no attribute 'tensor'
【问题讨论】:
标签: label classification spacy predict