Danielhome/cubeiro
0
1# Importações necessárias2from transformers import BertTokenizer, BertForSequenceClassification, RobertaTokenizer, RobertaForSequenceClassification3import torch4 5# Definir o modelo desejado: 'bert' para FinBERT ou 'roberta' para RoBERTa base6model_type = "bert" # Troque para 'roberta' conforme necessário7 8# Função para configurar o modelo e o tokenizador com base no tipo selecionado9def get_model_and_tokenizer(model_type, model_path, num_labels):10 if model_type == "bert":11 tokenizer = BertTokenizer.from_pretrained(model_path)12 model = BertForSequenceClassification.from_pretrained(model_path, num_labels=num_labels, ignore_mismatched_sizes=True)13 elif model_type == "roberta":14 tokenizer = RobertaTokenizer.from_pretrained(model_path)15 model = RobertaForSequenceClassification.from_pretrained(model_path, num_labels=num_labels)16 return tokenizer, model17 18# Caminho do modelo e número de classes predefinidas no FinBERT19finbert_path = "yiyanghkust/finbert-tone" # Caminho do modelo FinBERT20num_labels = 3 # FinBERT foi treinado com três classes para análise de tom financeiro21 22# Configurações do modelo e tokenizador para teste23tokenizer, model = get_model_and_tokenizer(model_type, finbert_path, num_labels=num_labels)24 25# Função de preparação dos dados (sem treinamento incremental)26def prepare_data(texts, tokenizer, max_length=128):27 encodings = tokenizer(texts, truncation=True, padding=True, max_length=max_length, return_tensors="pt")28 return encodings29 30# Exemplo de dados de entrada (você pode testar textos de conteúdo financeiro, por exemplo)31texts = ["O desempenho da empresa foi excelente este trimestre", 32 "Previsão de queda nas ações devido à instabilidade econômica",33 "Os resultados foram medianos e sem grandes oscilações"]34 35# Preparação dos dados usando o tokenizador selecionado36encodings = prepare_data(texts, tokenizer)37 38# Inferência com o modelo sem treinamento adicional39with torch.no_grad():40 outputs = model(**encodings)41 predictions = torch.argmax(outputs.logits, dim=-1)42 print("Predictions (Sentiment Classes):", predictions)43 