You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Hugging Face DistilBERT模型在MNLI验证集上准确率极低求助

问题描述

使用Hugging Face DistilBERT模型(基于MNLI微调的预训练模型)评估MNLI验证集时,仅获得7.90%的极低准确率,怀疑训练或数据预处理流程存在错误,恳请协助验证分词、数据加载及评估方法是否正确,相关代码如下:

from transformers import AutoTokenizer
from transformers import AutoModelForSequenceClassification
from datasets import load_dataset
import torch
from torch.utils.data import DataLoader
from tqdm import tqdm

# Mapping Numeric labels for text labels for MNLI
label_map = {0: "entailment", 1: "neutral", 2: "contradiction"}

# Tokenizes the MNLI dataset for the BERT model.
def prepare_data(tokenizer, dataset, max_length=512):
    # Tokenizes each pair of premise and hypothesis
    def tokenize_function(example):
        return tokenizer(
            example["premise"],  # The premise in the input text
            example["hypothesis"],  # The hypothesis in the input text
            truncation=True,  # Truncate sequences longer than max_length
            padding="max_length",  # Pad shorter sequences to max_length
            max_length=max_length  # Maximum token length for each input
        )
    
    # Apply tokenization to the entire dataset
    tokenized_dataset = dataset.map(tokenize_function, batched=True)
    
    # Debug: prints tokenization sample
    print("Sample tokenized example:", tokenized_dataset[0])
    
    # Checks sequence lengths
    for example in tokenized_dataset:
        input_length = len(example['input_ids'])  # Length of tokenized input sequence
        # Assert ensures that no sequence exceeds the defined max length
        assert input_length <= max_length, f"Input sequence exceeds max length: {input_length}"
    print("All tokenized sequences are within the max length.")
    return tokenized_dataset

# Computes the accuracy of predictions compared to references
def compute_accuracy(predictions, references):
    correct = sum(p == r for p, r in zip(predictions, references))  # Count correct predictions
    return correct / len(references)  # Return accuracy as a fraction

def main():
    # Loads pre-trained tokenizer and model
    model_name = "huggingface/distilbert-base-uncased-finetuned-mnli"
    tokenizer = AutoTokenizer.from_pretrained(model_name)  # Load tokenizer for the specified model
    model = AutoModelForSequenceClassification.from_pretrained(model_name)  # Load pre-trained sequence classification model
    model.eval()  # Set model to evaluation mode

    # Load MNLI dataset
    mnli_dataset = load_dataset("glue", "mnli")  # Load GLUE MNLI dataset
    test_data = mnli_dataset["validation_matched"]  # Use matched validation set for evaluation

    # Tokenize validation dataset
    tokenized_test_data = prepare_data(tokenizer, test_data)  # Tokenize the dataset

    # Convert dataset to PyTorch DataLoader
    def collate_fn(batch):
        # Define the keys to extract for model inputs
        keys = ["input_ids", "attention_mask"]
        # Create tensors for inputs and labels
        inputs = {key: torch.tensor([example[key] for example in batch]) for key in keys}
        labels = torch.tensor([example["label"] for example in batch])
        return inputs, labels

    test_loader = DataLoader(
        tokenized_test_data,  # Pass tokenized dataset
        batch_size=32,  # Adjust batch size for CPU
        collate_fn=collate_fn  # Collate function for preparing batches
    )

    # Evaluate the model
    predictions = []  # Store model predictions
    references = []  # Store ground truth labels
    for batch in tqdm(test_loader, desc="Evaluating"):  # Iterate over DataLoader
        inputs, labels = batch  # Extract inputs and labels from batch
        with torch.no_grad():  # Disable gradient computation for evaluation
            outputs = model(**inputs)  # Forward pass through the model
            logits = outputs.logits  # Extract logits from model outputs
            batch_predictions = torch.argmax(logits, dim=-1).tolist()  # Get predictions from logits
            predictions.extend(batch_predictions)  # Append predictions to the list
            references.extend(labels.tolist())  # Append references to the list

    # Debugging outputs
    print("Sample predictions (readable):", [label_map[p] for p in predictions[:5]])  # Print first 5 predictions
    print("Sample references (readable):", [label_map[r] for r in references[:5]])  # Print first 5 references
    
    # Compute and print accuracy
    accuracy = compute_accuracy(predictions, references)  # Calculate accuracy
    print(f"Accuracy on MNLI validation set: {accuracy * 100:.2f}%")  # Print accuracy as a percentage

if __name__ == "__main__":
    main()  # Run the main function
问题排查与修正

1. 核心错误:模型名称不正确

你使用的模型ID huggingface/distilbert-base-uncased-finetuned-mnli 是错误的,正确的预训练模型ID为 distilbert-base-uncased-finetuned-mnli。错误的模型名会导致加载到非目标模型(或权重不匹配的模型),这是准确率极低的直接原因。

2. 数据预处理优化(非核心问题,但更规范)

在prepare_data函数末尾添加数据集格式转换,避免手动写collate_fn的潜在问题:

tokenized_dataset.set_format("torch", columns=["input_ids", "attention_mask", "label"])

之后DataLoader可以直接使用默认的collate_fn,无需自定义。

3. 验证其他流程

  • 分词逻辑正确:同时传入premise和hypothesis,设置了截断和padding到最大长度,且通过断言验证了序列长度,这部分没有问题。
  • 评估逻辑正确:compute_accuracy的计算方式符合分类任务准确率的定义,标签映射也与MNLI数据集及预训练模型的标签定义一致。
修正后的代码
from transformers import AutoTokenizer
from transformers import AutoModelForSequenceClassification
from datasets import load_dataset
import torch
from torch.utils.data import DataLoader
from tqdm import tqdm

# Mapping Numeric labels for text labels for MNLI
label_map = {0: "entailment", 1: "neutral", 2: "contradiction"}

# Tokenizes the MNLI dataset for the BERT model.
def prepare_data(tokenizer, dataset, max_length=512):
    # Tokenizes each pair of premise and hypothesis
    def tokenize_function(example):
        return tokenizer(
            example["premise"],
            example["hypothesis"],
            truncation=True,
            padding="max_length",
            max_length=max_length
        )
    
    tokenized_dataset = dataset.map(tokenize_function, batched=True)
    
    # 设置数据集格式为PyTorch张量
    tokenized_dataset.set_format("torch", columns=["input_ids", "attention_mask", "label"])
    
    # Debug: 打印样本
    print("Sample tokenized example:", tokenized_dataset[0])
    
    # 验证序列长度
    for example in tokenized_dataset:
        input_length = len(example['input_ids'])
        assert input_length <= max_length, f"Input sequence exceeds max length: {input_length}"
    print("All tokenized sequences are within the max length.")
    return tokenized_dataset

# Computes the accuracy of predictions compared to references
def compute_accuracy(predictions, references):
    correct = sum(p == r for p, r in zip(predictions, references))
    return correct / len(references)

def main():
    # 加载正确的预训练模型和分词器
    model_name = "distilbert-base-uncased-finetuned-mnli"
    tokenizer = AutoTokenizer.from_pretrained(model_name)
    model = AutoModelForSequenceClassification.from_pretrained(model_name)
    model.eval()

    # Load MNLI dataset
    mnli_dataset = load_dataset("glue", "mnli")
    test_data = mnli_dataset["validation_matched"]

    # Tokenize validation dataset
    tokenized_test_data = prepare_data(tokenizer, test_data)

    # 使用默认collate_fn的DataLoader
    test_loader = DataLoader(
        tokenized_test_data,
        batch_size=32,
        shuffle=False
    )

    # Evaluate the model
    predictions = []
    references = []
    for batch in tqdm(test_loader, desc="Evaluating"):
        inputs = {k: batch[k] for k in ["input_ids", "attention_mask"]}
        labels = batch["label"]
        with torch.no_grad():
            outputs = model(**inputs)
            logits = outputs.logits
            batch_predictions = torch.argmax(logits, dim=-1).tolist()
            predictions.extend(batch_predictions)
            references.extend(labels.tolist())

    # Debugging outputs
    print("Sample predictions (readable):", [label_map[p] for p in predictions[:5]])
    print("Sample references (readable):", [label_map[r] for r in references[:5]])
    
    # Compute and print accuracy
    accuracy = compute_accuracy(predictions, references)
    print(f"Accuracy on MNLI validation set: {accuracy * 100:.2f}%")

if __name__ == "__main__":
    main()

内容的提问来源于stack exchange,提问作者IAmNewToThis

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.15 15:37:06