Reasoning-Critical Span Scorer

DistilBERT-based binary classifier predicting whether a span in a reasoning question (GSM8K / BBH) is counterfactually critical -- i.e. masking it changes a correct answer to incorrect. Trained on consensus labels from 3 LLMs via masking-based counterfactual testing.

5-fold cross-validation ensemble; each foldN/scorer.pt is a full model state_dict trained with problems in that fold held out. split.json maps problem uid -> fold index.

Architecture: distilbert-base-uncased encoder, mean-pooled over the span's tokens, single linear classification head.

This is NOT a standard HF model -- it needs the custom SpanScorer class to load. See the Hugging Face repository for train_scorer.py with the class definition, or copy it below. This model uses PyTorch only and does not require TensorFlow.

# minimal SpanScorer class needed to load the weights
import torch
from torch import nn
from transformers import DistilBertModel

class SpanScorer(nn.Module):
    def __init__(self, encoder_name="distilbert-base-uncased", dropout=0.1):
        super().__init__()
        self.encoder = DistilBertModel.from_pretrained(encoder_name)
        self.dropout = nn.Dropout(dropout)
        self.classifier = nn.Linear(self.encoder.config.hidden_size, 1)

    def forward(self, input_ids, attention_mask, span_masks):
        h = self.encoder(input_ids=input_ids, attention_mask=attention_mask).last_hidden_state
        num = torch.einsum("bsl,blh->bsh", span_masks, h)
        den = span_masks.sum(-1, keepdim=True).clamp(min=1.0)
        return self.classifier(self.dropout(num / den)).squeeze(-1)

model = SpanScorer()
model.load_state_dict(torch.load("fold0/scorer.pt", map_location="cpu"))
model.eval()

Trained for the ANLP project "Reasoning-Critical Prompt Compression" (Team Symbiote).

Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support