decision_maker / decision_model.py
HrushikeshGangane's picture
Publish decision_maker inference bundle
c800ccd verified
Raw History Blame Contribute Delete
3.15 kB
"""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)