"""Bidirectional decision encoder with query/key option scoring.""" from __future__ import annotations from dataclasses import asdict, dataclass from typing import Any import torch from torch import Tensor, nn from transformers import AutoModel, AutoTokenizer SPECIAL_TOKENS = ["[STATE]", "[QUESTION]", "[OPTION]", "[DECIDE]"] @dataclass class ModelConfig: encoder_name: str = "bert-base-uncased" decision_layers: int = 4 decision_heads: int = 16 decision_ffn_dim: int = 4096 dropout: float = 0.1 score_dim: int = 256 def to_dict(self) -> dict[str, Any]: return asdict(self) def build_tokenizer(encoder_name: str): tokenizer = AutoTokenizer.from_pretrained(encoder_name, use_fast=True) tokenizer.add_special_tokens({"additional_special_tokens": SPECIAL_TOKENS}) return tokenizer class DecisionModel(nn.Module): def __init__(self, config: ModelConfig): super().__init__() self.config = config self.encoder = AutoModel.from_pretrained(config.encoder_name) hidden = self.encoder.config.hidden_size if hidden % config.decision_heads: raise ValueError(f"hidden_size={hidden} is not divisible by decision_heads={config.decision_heads}") layer = nn.TransformerEncoderLayer( d_model=hidden, nhead=config.decision_heads, dim_feedforward=config.decision_ffn_dim, dropout=config.dropout, activation="gelu", batch_first=True, norm_first=True, ) self.decision_transformer = nn.TransformerEncoder( layer, num_layers=config.decision_layers, enable_nested_tensor=False ) self.final_norm = nn.LayerNorm(hidden) self.query = nn.Linear(hidden, config.score_dim, bias=False) self.key = nn.Linear(hidden, config.score_dim, bias=False) self.scale = config.score_dim ** -0.5 def resize_token_embeddings(self, vocab_size: int) -> None: self.encoder.resize_token_embeddings(vocab_size) def forward( self, input_ids: Tensor, attention_mask: Tensor, decide_positions: Tensor, option_positions: Tensor, option_mask: Tensor, ) -> Tensor: """Return one unnormalized logit per real option, padded slots = -inf.""" base = self.encoder(input_ids=input_ids, attention_mask=attention_mask).last_hidden_state # Transformer expects True where tokens should be ignored. contextual = self.decision_transformer(base, src_key_padding_mask=~attention_mask.bool()) contextual = self.final_norm(contextual) batch = torch.arange(input_ids.size(0), device=input_ids.device) decision_vec = contextual[batch, decide_positions] # [B, H] safe_option_positions = option_positions.clamp_min(0) option_vecs = contextual[batch[:, None], safe_option_positions] # [B, N, H] q = self.query(decision_vec).unsqueeze(1) k = self.key(option_vecs) logits = (q * k).sum(-1) * self.scale return logits.masked_fill(~option_mask, torch.finfo(logits.dtype).min)