【问题标题】:How to use a batch size bigger than zero in Bert Sequence Classification如何在 Bert 序列分类中使用大于零的批量大小
【发布时间】:2020-05-26 22:19:57
【问题描述】:

Hugging Face documentation describes如何使用Bert模型进行序列分类:

from transformers import BertTokenizer, BertForSequenceClassification
import torch

tokenizer = BertTokenizer.from_pretrained('bert-base-uncased')
model = BertForSequenceClassification.from_pretrained('bert-base-uncased')

input_ids = torch.tensor(tokenizer.encode("Hello, my dog is cute", add_special_tokens=True)).unsqueeze(0)  # Batch size 1
labels = torch.tensor([1]).unsqueeze(0)  # Batch size 1
outputs = model(input_ids, labels=labels)

loss, logits = outputs[:2]

但是,只有批量大小 1 的示例。当我们有一个短语列表并想要使用更大的批量大小时,如何实现它?

【问题讨论】:

    标签: python huggingface-transformers


    【解决方案1】:

    在该示例中,unsqueeze 用于向输入/标签添加维度,因此它是大小为 (batch_size, sequence_length) 的数组。如果您想使用大于 1 的批量大小,则可以构建一个序列数组,如下例所示:

    from transformers import BertTokenizer, BertForSequenceClassification
    import torch
    
    tokenizer = BertTokenizer.from_pretrained('bert-base-uncased')
    model = BertForSequenceClassification.from_pretrained('bert-base-uncased')
    
    sequences = ["Hello, my dog is cute", "My dog is cute as well"]
    input_ids = torch.tensor([tokenizer.encode(sequence, add_special_tokens=True) for sequence in sequences])
    labels = torch.tensor([[1], [0]]) # Labels depend on the task
    outputs = model(input_ids, labels=labels)
    
    loss, logits = outputs[:2]
    

    在该示例中,两个序列都以相同数量的标记进行编码,因此很容易构建包含两个序列的张量,但如果它们具有不同数量的元素,则需要填充序列并告诉模型哪些标记它应该注意(以便它忽略填充值)使用注意掩码。

    glossary 中有一个关于注意面具的条目,解释了它们的目的和用途。您在调用模型的 forward 方法时将此注意掩码传递给模型。

    【讨论】:

    • 成功了,谢谢!正如你所说,这种情况是巧合,因为两个短语具有相同数量的标记,所以我实现了一个函数来在不同长度的短语中添加一个 0 填充,但我仍然必须使用注意力掩码来实现。至于标签,我怎么知道什么时候使用[1]、[0]或其他单元素数组?
    • 这取决于您的任务。您需要为序列分类确定不同的类别,每个类别都有一个数字索引(例如,第一类为 0,第二类为 1,第三类为 2)。然后,您必须根据您的特定任务微调此模型。
    • 所以,如果我的问题有 20 个类,并且我试图预测 3 个短语,我应该构建一个 [[0, 1, ..., 19], [0, 1, . .., 19], [0, 1, ..., 19]] 数组?
    猜你喜欢
    • 1970-01-01
    • 2018-02-21
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2021-01-13
    • 2018-06-27
    • 1970-01-01
    • 2017-08-05
    相关资源
    最近更新 更多