compliantLLM / modeling_compliant_llm.py
Martin Navrátil
Upload 11 files
4689b4d verified
Raw
History Blame Contribute Delete
3.4 kB
"""Hugging Face model implementation for compliantLLM inference."""
from dataclasses import dataclass
from typing import Optional
import torch
from torch import Tensor, nn
from transformers import PreTrainedModel
from transformers.utils import ModelOutput
from .configuration_compliant_llm import CompliantLLMConfig
@dataclass
class CompliantLLMOutput(ModelOutput):
logits: Optional[Tensor] = None
class CompliantLLMModel(PreTrainedModel):
config_class = CompliantLLMConfig
base_model_prefix = "compliant_llm"
main_input_name = "input_ids"
def __init__(self, config):
super().__init__(config)
self.token_embedding = nn.Embedding(config.input_vocab_size, config.d_model)
self.position_embedding = nn.Embedding(config.max_context, config.d_model)
layer = nn.TransformerEncoderLayer(
d_model=config.d_model,
nhead=config.n_heads,
dim_feedforward=config.ffn_dim,
dropout=config.dropout,
activation="gelu",
batch_first=True,
norm_first=True,
)
self.encoder = nn.TransformerEncoder(
layer,
num_layers=config.n_layers,
enable_nested_tensor=False,
)
self.output_positions = nn.Parameter(torch.empty(config.output_length, config.d_model))
self.output_norm = nn.LayerNorm(config.d_model)
self.output_head = nn.Linear(config.d_model, config.output_vocab_size)
self.post_init()
def forward(self, input_ids, attention_mask=None, **kwargs):
del kwargs
if input_ids.ndim != 2:
raise ValueError("input_ids must have shape [batch, sequence]")
_, sequence_length = input_ids.shape
if sequence_length > self.config.max_context:
raise ValueError(f"sequence exceeds {self.config.max_context}-token context")
if attention_mask is None:
attention_mask = torch.ones_like(input_ids, dtype=torch.bool)
positions = torch.arange(sequence_length, device=input_ids.device)
hidden = self.token_embedding(input_ids)
hidden = hidden + self.position_embedding(positions)[None, :, :]
hidden = self.encoder(hidden, src_key_padding_mask=~attention_mask.bool())
weights = attention_mask.to(hidden.dtype).unsqueeze(-1)
pooled = (hidden * weights).sum(dim=1) / weights.sum(dim=1).clamp_min(1.0)
output_hidden = pooled[:, None, :] + self.output_positions[None, :, :]
logits = self.output_head(self.output_norm(output_hidden))
return CompliantLLMOutput(logits=logits)
@torch.inference_mode()
def generate(self, input_ids, attention_mask=None, **kwargs):
"""Return the three output-vocabulary IDs; generation is non-autoregressive."""
del kwargs
return self(input_ids=input_ids, attention_mask=attention_mask).logits.argmax(dim=-1)
def decode_output(self, output_ids):
"""Decode one generated sequence from the separate output vocabulary."""
if isinstance(output_ids, Tensor):
output_ids = output_ids.detach().cpu().tolist()
if any(token < 0 or token >= self.config.output_vocab_size for token in output_ids):
raise ValueError("output token ID outside the three-token vocabulary")
return "".join(self.config.output_tokens[token] for token in output_ids)