Download decision_model.py from HrushikeshGangane/decision_maker: direct link, hf CLI and curl.
- Browser
- Download file 3.15 kB
-
https://huggingface.co/HrushikeshGangane/decision_maker/resolve/main/decision_model.py
- Command line
-
hf download hf://HrushikeshGangane/decision_maker/decision_model.py
-
curl -L -o decision_model.py https://huggingface.co/HrushikeshGangane/decision_maker/resolve/main/decision_model.py
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]"] | |
| 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) | |