Pivot / modeling_pivot.py
Q1z's picture
Pivot H200 full retrain
14bf8c2 verified
Raw History Blame Contribute Delete
5.57 kB
"""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(),
}