WaterSheep / modeling_watersheep.py
samratduttaofficial's picture
Initial commit
07f1ef2
Raw History Blame Contribute Delete
2.43 kB
from dataclasses import dataclass
import torch
from torch import nn
from transformers import AutoConfig, AutoModel, PreTrainedConfig, PreTrainedModel
from transformers.utils import ModelOutput
class WaterSheepConfig(PreTrainedConfig):
model_type = "watersheep"
def __init__(self, encoder_config=None, head_layers=1, max_len=512, max_question_tokens=96,
max_option_tokens=32, max_options=10, temperatures=None, multi_threshold=0.5, **kwargs):
self.encoder_config = encoder_config or {}
self.head_layers = head_layers
self.max_len = max_len
self.max_question_tokens = max_question_tokens
self.max_option_tokens = max_option_tokens
self.max_options = max_options
self.temperatures = temperatures or {}
self.multi_threshold = multi_threshold
super().__init__(**kwargs)
@dataclass
class WaterSheepOutput(ModelOutput):
logits: torch.FloatTensor = None
class WaterSheepModel(PreTrainedModel):
config_class = WaterSheepConfig
base_model_prefix = "watersheep"
def __init__(self, config):
super().__init__(config)
enc = AutoConfig.for_model(**config.encoder_config)
if hasattr(enc, "reference_compile"):
enc.reference_compile = False
self.encoder = AutoModel.from_config(enc)
h = enc.hidden_size
self.proj = nn.Sequential(nn.Dropout(0.0), nn.Linear(h, h), nn.GELU(), nn.LayerNorm(h))
self.mix = None
if config.head_layers > 0:
layer = nn.TransformerEncoderLayer(h, nhead=max(1, h // 64), dim_feedforward=2 * h, dropout=0.0,
activation="gelu", batch_first=True, norm_first=True)
self.mix = nn.TransformerEncoder(layer, config.head_layers, enable_nested_tensor=False)
self.out = nn.Linear(h, 1)
self.post_init()
def forward(self, input_ids, attention_mask, option_positions, option_mask, **kwargs):
hs = self.encoder(input_ids=input_ids, attention_mask=attention_mask).last_hidden_state
idx = option_positions.clamp(min=0).unsqueeze(-1).expand(-1, -1, hs.size(-1))
x = self.proj(torch.gather(hs, 1, idx))
if self.mix is not None:
x = self.mix(x, src_key_padding_mask=~option_mask)
logits = self.out(x).squeeze(-1).float()
return WaterSheepOutput(logits=logits.masked_fill(~option_mask, -1e4))