"""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)