"""Self-contained Hugging Face runtime for Pivot.""" import torch from torch import nn from transformers import AutoConfig, PreTrainedModel from .configuration_pivot import PivotConfig from .modeling_lfm2_bidirectional import Lfm2BidirectionalModel from .pivot_model import HuggingFaceEncoder, MLPScorer, SetBrierEncoder, StubEncoder from .pivot_infer import decide_native, decide_typed, predict class PivotModel(PreTrainedModel): config_class = PivotConfig base_model_prefix = "network" _no_split_modules = ["SetBrierEncoder"] def __init__(self, config): super().__init__(config) cfg = config.dsbt_config backbone = cfg["backbone"] if backbone.get("encoder") == "stub": encoder = StubEncoder( int(backbone.get("stub_vocab_size", 256)), int(backbone.get("stub_hidden_size", 32)), ) else: body_config = dict(config.encoder_config) model_type = body_config.pop("model_type") body = Lfm2BidirectionalModel(AutoConfig.for_model(model_type, **body_config)) encoder = HuggingFaceEncoder.__new__(HuggingFaceEncoder) nn.Module.__init__(encoder) encoder.model = body encoder.hidden_size = int(body.config.hidden_size) encoder.pretrained_audit = {"packaged_runtime": True} hidden = encoder.hidden_size scorer_cfg = cfg["scorer"] if scorer_cfg.get("type", "mlp") != "mlp": raise ValueError("This Pivot package supports the reviewed MLP scorer only") scorer = MLPScorer(hidden, int(scorer_cfg.get("hidden_size", hidden))) self.network = SetBrierEncoder(encoder, scorer) self.post_init() def forward(self, ctx_ids, ctx_mask, opt_ids, opt_mask, opt_attn): return self.network(ctx_ids, ctx_mask, opt_ids, opt_mask, opt_attn) def _limits(self): cfg = self.config.dsbt_config data = cfg["data"] serving = cfg.get("serving") or {} return ( int(serving.get("max_context_tokens", data["max_context_tokens"])), int(serving.get("max_option_tokens", data["max_option_tokens"])), ) @torch.no_grad() def decide(self, tokenizer, state, questions): self.eval() max_context_tokens, max_option_tokens = self._limits() return decide_typed( self.network, tokenizer, state, questions, max_context_tokens=max_context_tokens, max_option_tokens=max_option_tokens, device=next(self.parameters()).device, model_id="Pivot", ) @torch.no_grad() def decide_native(self, tokenizer, context, candidates): self.eval() max_context_tokens, max_option_tokens = self._limits() return decide_native( self.network, tokenizer, context, candidates, max_context_tokens=max_context_tokens, max_option_tokens=max_option_tokens, device=next(self.parameters()).device, ) @torch.no_grad() def choose(self, tokenizer, context, options): self.eval() max_context_tokens, max_option_tokens = self._limits() return predict( self.network, tokenizer, context, list(options), max_context_tokens=max_context_tokens, max_option_tokens=max_option_tokens, device=next(self.parameters()).device, ) @torch.no_grad() def _encode_texts(self, tokenizer, texts, *, max_length, padding): batch = tokenizer( list(texts), max_length=int(max_length), padding=padding, truncation=True, return_tensors="pt", ) device = next(self.parameters()).device return self.network.encode( batch["input_ids"].to(device), batch["attention_mask"].to(device), ) @torch.no_grad() def encode_context(self, tokenizer, context): self.eval() max_context_tokens, _ = self._limits() return self._encode_texts( tokenizer, [context], max_length=max_context_tokens, padding=True, )[0] @torch.no_grad() def encode_candidates(self, tokenizer, candidate_texts): self.eval() _, max_option_tokens = self._limits() return self._encode_texts( tokenizer, list(candidate_texts), max_length=max_option_tokens, padding="max_length", ) @torch.no_grad() def choose_cached(self, context_embedding, candidate_embeddings, candidate_texts): self.eval() texts = list(candidate_texts) if candidate_embeddings.ndim == 2: k = int(candidate_embeddings.shape[0]) elif candidate_embeddings.ndim == 3 and candidate_embeddings.shape[0] == 1: k = int(candidate_embeddings.shape[1]) else: raise ValueError("candidate_embeddings must be [K,d] or [1,K,d]") if len(texts) != k or k < 2: raise ValueError("candidate embedding/text mismatch") out = self.network.score_preencoded(context_embedding, candidate_embeddings) probs = out.probs[0].detach().cpu().float() index = int(out.pred_index[0].item()) return { "choice": texts[index], "index": index, "probs": probs.tolist(), }