Text Classification
Transformers
Safetensors
lfm2
feature-extraction
betterwright
accessibility
browser-agent
reranking
long-context
custom_code
Instructions to use ProCreations/betterwright-encoder-350m with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use ProCreations/betterwright-encoder-350m with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-classification", model="ProCreations/betterwright-encoder-350m", trust_remote_code=True)# Load model directly from transformers import AutoTokenizer, AutoModel tokenizer = AutoTokenizer.from_pretrained("ProCreations/betterwright-encoder-350m", trust_remote_code=True) model = AutoModel.from_pretrained("ProCreations/betterwright-encoder-350m", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
| """BetterWright task-conditioned relevance heads on LFM2.5 Encoder.""" | |
| from __future__ import annotations | |
| from dataclasses import dataclass | |
| from typing import Optional | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from transformers.modeling_outputs import ModelOutput | |
| from transformers.models.lfm2.configuration_lfm2 import Lfm2Config | |
| from transformers.models.lfm2.modeling_lfm2 import Lfm2PreTrainedModel | |
| from .modeling_lfm2_bidirectional import Lfm2BidirectionalModel, _install_patches | |
| class BetterWrightEncoderOutput(ModelOutput): | |
| loss: Optional[torch.Tensor] = None | |
| logits: Optional[torch.Tensor] = None | |
| token_logits: Optional[torch.Tensor] = None | |
| uncertainty_logits: Optional[torch.Tensor] = None | |
| token_loss: Optional[torch.Tensor] = None | |
| token_rank_loss: Optional[torch.Tensor] = None | |
| last_hidden_state: Optional[torch.Tensor] = None | |
| class BetterWrightEncoder(Lfm2PreTrainedModel): | |
| """Cross-encoder chunk scorer plus per-token relevance and fallback heads.""" | |
| config_class = Lfm2Config | |
| base_model_prefix = "lfm2" | |
| def __init__(self, config: Lfm2Config): | |
| _install_patches() | |
| config.use_cache = False | |
| super().__init__(config) | |
| self.lfm2 = Lfm2BidirectionalModel(config) | |
| self.relevance_head = nn.Sequential( | |
| nn.Linear(config.hidden_size, config.hidden_size // 2), | |
| nn.SiLU(), | |
| nn.Dropout(0.05), | |
| nn.Linear(config.hidden_size // 2, 1), | |
| ) | |
| self.token_head = nn.Linear(config.hidden_size, 1) | |
| self.uncertainty_head = nn.Linear(config.hidden_size, 1) | |
| self.post_init() | |
| def get_input_embeddings(self): | |
| return self.lfm2.embed_tokens | |
| def set_input_embeddings(self, value): | |
| self.lfm2.embed_tokens = value | |
| def forward( | |
| self, | |
| input_ids: Optional[torch.LongTensor] = None, | |
| attention_mask: Optional[torch.Tensor] = None, | |
| position_ids: Optional[torch.LongTensor] = None, | |
| inputs_embeds: Optional[torch.FloatTensor] = None, | |
| labels: Optional[torch.Tensor] = None, | |
| token_labels: Optional[torch.Tensor] = None, | |
| return_dict: bool = True, | |
| **kwargs, | |
| ) -> BetterWrightEncoderOutput | tuple: | |
| outputs = self.lfm2( | |
| input_ids=input_ids, | |
| attention_mask=attention_mask, | |
| position_ids=position_ids, | |
| inputs_embeds=inputs_embeds, | |
| use_cache=False, | |
| return_dict=True, | |
| **kwargs, | |
| ) | |
| hidden = outputs.last_hidden_state | |
| token_logits = self.token_head(hidden).squeeze(-1) | |
| if attention_mask is None: | |
| pooling_weights = torch.softmax(token_logits, dim=1).unsqueeze(-1) | |
| else: | |
| masked_pool_logits = token_logits.masked_fill(attention_mask == 0, -1e4) | |
| pooling_weights = torch.softmax(masked_pool_logits, dim=1).unsqueeze(-1) | |
| # Learned token pooling lets pair-level relevance supervision teach the | |
| # token head where evidence lives before the exact-span curriculum. | |
| pooled = (hidden * pooling_weights.to(hidden.dtype)).sum(dim=1) | |
| logits = self.relevance_head(pooled).squeeze(-1) | |
| uncertainty_logits = self.uncertainty_head(pooled).squeeze(-1) | |
| loss = None | |
| token_loss = None | |
| token_rank_loss = None | |
| if labels is not None: | |
| labels = labels.to(logits.dtype) | |
| positive_weight = float(getattr(self.config, "betterwright_positive_weight", 3.0)) | |
| relevance_loss = F.binary_cross_entropy_with_logits( | |
| logits, | |
| labels, | |
| pos_weight=torch.tensor(positive_weight, device=logits.device, dtype=logits.dtype), | |
| ) | |
| # Predict disagreement/ambiguity most strongly around the decision boundary. | |
| uncertainty_target = 1.0 - (labels - torch.sigmoid(logits).detach()).abs() | |
| uncertainty_loss = F.binary_cross_entropy_with_logits( | |
| uncertainty_logits, uncertainty_target.clamp(0, 1) | |
| ) | |
| loss = relevance_loss + 0.08 * uncertainty_loss | |
| if token_labels is not None: | |
| valid = token_labels >= 0 | |
| if valid.any(): | |
| valid_labels = token_labels[valid] | |
| positive_count = (valid_labels == 1).sum() | |
| negative_count = (valid_labels == 0).sum() | |
| if positive_count: | |
| # Required lines are exceptionally sparse in a 64K tree. | |
| # Batch-adaptive weighting prevents the all-negative token | |
| # solution while remaining capped for stable BF16 updates. | |
| token_positive_weight = float( | |
| (negative_count / positive_count).clamp(20, 512) | |
| ) | |
| else: | |
| token_positive_weight = 1.0 | |
| token_loss = F.binary_cross_entropy_with_logits( | |
| token_logits[valid], | |
| valid_labels.to(token_logits.dtype), | |
| pos_weight=torch.tensor( | |
| token_positive_weight, | |
| device=token_logits.device, | |
| dtype=token_logits.dtype, | |
| ), | |
| ) | |
| # Runtime pruning ranks lines by their highest-scoring token. | |
| # BCE alone can achieve a reasonable average while leaving a | |
| # few distractor tokens above the evidence. Explicitly compare | |
| # each positive row against its hardest negative tokens so the | |
| # training objective matches that production ranking behavior. | |
| token_rank_losses = [] | |
| for row_index in range(token_logits.shape[0]): | |
| row_valid = valid[row_index] | |
| row_labels = token_labels[row_index] | |
| positives = token_logits[row_index][row_valid & (row_labels == 1)] | |
| negatives = token_logits[row_index][row_valid & (row_labels == 0)] | |
| if positives.numel() == 0 or negatives.numel() == 0: | |
| continue | |
| hard_count = min(64, negatives.numel()) | |
| hard_negatives = torch.topk( | |
| negatives.float(), hard_count, sorted=False | |
| ).values | |
| token_rank_losses.append( | |
| F.softplus( | |
| 1.0 + hard_negatives.mean() - positives.float().mean() | |
| ) | |
| ) | |
| if token_rank_losses: | |
| token_rank_loss = torch.stack(token_rank_losses).mean() | |
| else: | |
| token_rank_loss = token_loss.new_zeros(()) | |
| token_objective = 0.75 * token_loss + token_rank_loss | |
| loss = token_objective if loss is None else loss + token_objective | |
| result = BetterWrightEncoderOutput( | |
| loss=loss, | |
| logits=logits, | |
| token_logits=token_logits, | |
| uncertainty_logits=uncertainty_logits, | |
| token_loss=token_loss, | |
| token_rank_loss=token_rank_loss, | |
| last_hidden_state=hidden, | |
| ) | |
| return result if return_dict else tuple(result.values()) | |