Sobes.tech
Middle — Senior

Como fazer fine-tuning do BERT para uma tarefa de classificação? O que pegar da saída?

sobes.tech IA

Resposta da IA

Para ajustar finamente o BERT para tarefas de classificação, geralmente usa-se a saída do token especial [CLS], que é utilizado pelo modelo para agregar informações sobre toda a sequência.

Passos principais:

  1. Adicionar uma camada de classificação sobre o BERT — geralmente uma camada linear que recebe o vetor da última camada do BERT correspondente ao token [CLS].

  2. Os dados de entrada são fornecidos ao BERT com o token [CLS] no início.

  3. Extrair a embedding [CLS] da última camada oculta do BERT. Este vetor, por exemplo, de tamanho 768 (para o modelo base), contém a informação condensada de toda a sequência.

  4. Enviar esse vetor ao classificador (por exemplo, linear + softmax para classificação multiclasse).

  5. Treinar todo o modelo (BERT + classificador) conjuntamente na tarefa alvo, usando uma função de perda adequada (por exemplo, CrossEntropyLoss).

Exemplo em PyTorch:

from transformers import BertModel, BertTokenizer
import torch.nn as nn

class BertClassifier(nn.Module):
    def __init__(self, bert_model_name, num_classes):
        super().__init__()
        self.bert = BertModel.from_pretrained(bert_model_name)
        self.classifier = nn.Linear(self.bert.config.hidden_size, num_classes)

    def forward(self, input_ids, attention_mask):
        outputs = self.bert(input_ids=input_ids, attention_mask=attention_mask)
        cls_output = outputs.last_hidden_state[:, 0, :]  # token [CLS]
        logits = self.classifier(cls_output)
        return logits

Portanto, para classificação, pega-se especificamente a saída do token [CLS] da última camada do BERT.