bert-agent / bert /model.py
SNAPKITTYWEST's picture
push from SNAPKITTYWEST/bert-agent
30f011f verified
Raw
History Blame Contribute Delete
3.54 kB
"""
bert/model.py
DeBERTa-v3 Cross-Encoder for entailment verification.
Architecture: Cross-Encoder (NOT Bi-Encoder).
Both premise and hypothesis are concatenated and fed through the transformer
together so self-attention can directly compare entities, negations, and dates
across the premise-hypothesis boundary. This is critical for catching
hallucinations like flipped dates or switched subjects.
Labels: 0=Contradiction, 1=Neutral, 2=Entailment
"""
from __future__ import annotations
import torch
import torch.nn as nn
from dataclasses import dataclass
from typing import Optional
from transformers import AutoModelForSequenceClassification, AutoTokenizer, PreTrainedModel
@dataclass
class CrossEncoderConfig:
backbone: str = "microsoft/deberta-v3-base"
num_labels: int = 3 # 0=Contradiction 1=Neutral 2=Entailment
dropout: float = 0.1
max_length: int = 512
# Class weights: penalise false-positive Entailment heavily.
# Contradiction=2.0, Neutral=1.5, Entailment=1.0
class_weights: tuple = (2.0, 1.5, 1.0)
class BertCrossEncoderVerifier(nn.Module):
"""
DeBERTa-v3 Cross-Encoder entailment verifier.
DeBERTa is chosen over BERT/RoBERTa for its disentangled attention
mechanism which handles positional reasoning significantly better —
critical when LLM-generated claims reorder events from the source chunk.
"""
def __init__(self, config: CrossEncoderConfig) -> None:
super().__init__()
self.config = config
self.model: PreTrainedModel = AutoModelForSequenceClassification.from_pretrained(
config.backbone,
num_labels=config.num_labels,
hidden_dropout_prob=config.dropout,
attention_probs_dropout_prob=config.dropout,
)
def forward(
self,
input_ids: torch.Tensor,
attention_mask: torch.Tensor,
token_type_ids: Optional[torch.Tensor] = None,
labels: Optional[torch.Tensor] = None,
) -> dict:
outputs = self.model(
input_ids=input_ids,
attention_mask=attention_mask,
token_type_ids=token_type_ids,
labels=None, # compute loss manually with class weights
)
logits = outputs.logits # (batch, 3)
result = {"logits": logits}
if labels is not None:
weights = torch.tensor(
self.config.class_weights,
dtype=torch.float,
device=logits.device,
)
loss_fn = nn.CrossEntropyLoss(weight=weights)
result["loss"] = loss_fn(logits, labels)
return result
@torch.no_grad()
def predict_entailment_score(
self,
input_ids: torch.Tensor,
attention_mask: torch.Tensor,
token_type_ids: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""Return softmax probability of Entailment class (index 2)."""
out = self.forward(input_ids, attention_mask, token_type_ids)
probs = torch.softmax(out["logits"], dim=-1)
return probs[:, 2] # entailment column
def build_model(config: CrossEncoderConfig | None = None) -> BertCrossEncoderVerifier:
if config is None:
config = CrossEncoderConfig()
return BertCrossEncoderVerifier(config)
def load_tokenizer(backbone: str = "microsoft/deberta-v3-base") -> AutoTokenizer:
return AutoTokenizer.from_pretrained(backbone)