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).