oryn / inference.py
timosarkar's picture
added files
cac788c verified
Raw History Blame Contribute Delete
6.53 kB
import argparse
import json
import time
from pathlib import Path
import torch
import torch.nn as nn
import torch.nn.functional as F
from transformers import AutoConfig, AutoModel, AutoTokenizer
MODEL_DIR = Path(__file__).resolve().parent
# ---------------------------------------------------------------------
# Model architecture
# ---------------------------------------------------------------------
class DecisionHead(nn.Module):
def __init__(self, hidden_size):
super().__init__()
heads = next(
head_count
for head_count in (8, 4, 2, 1)
if hidden_size % head_count == 0
)
layer = nn.TransformerEncoderLayer(
d_model=hidden_size,
nhead=heads,
dim_feedforward=hidden_size * 4,
batch_first=True,
activation="gelu",
)
self.encoder = nn.TransformerEncoder(
layer,
num_layers=2,
enable_nested_tensor=False,
)
self.scorer = nn.Linear(hidden_size, 1)
def forward(self, hidden, attention_mask):
return self.encoder(
hidden,
src_key_padding_mask=attention_mask == 0,
)
class LayaModel(nn.Module):
def __init__(self, backbone_name):
super().__init__()
# Downloads only the backbone configuration when it is not cached.
backbone_config = AutoConfig.from_pretrained(
backbone_name
)
self.backbone = AutoModel.from_config(
backbone_config
)
self.head = DecisionHead(
backbone_config.hidden_size
)
def forward(
self,
input_ids,
attention_mask,
mask_positions,
):
hidden = self.backbone(
input_ids=input_ids,
attention_mask=attention_mask,
).last_hidden_state
hidden = self.head(
hidden,
attention_mask,
)
return [
self.head.scorer(
hidden[index, positions]
).squeeze(-1)
for index, positions in enumerate(mask_positions)
]
# ---------------------------------------------------------------------
# Load model
# ---------------------------------------------------------------------
with open(
MODEL_DIR / "config.json",
encoding="utf-8",
) as file:
settings = json.load(file)
BACKBONE = settings["backbone"]
MAX_LEN = settings.get("max_len", 256)
device = (
"cuda"
if torch.cuda.is_available()
else "mps"
if torch.backends.mps.is_available()
else "cpu"
)
tokenizer = AutoTokenizer.from_pretrained(
MODEL_DIR,
local_files_only=True,
)
model = LayaModel(BACKBONE)
state_dict = torch.load(
MODEL_DIR / "model.bin",
map_location="cpu",
weights_only=True,
)
model.load_state_dict(
state_dict,
strict=True,
)
model = model.to(device)
model.eval()
# ---------------------------------------------------------------------
# Inference
# ---------------------------------------------------------------------
@torch.inference_mode()
def predict(text, question, options):
option_text = ", ".join(
f"{option} {tokenizer.mask_token}"
for option in options
)
prefix = f"question: {question}\nstate: "
suffix = f"\noptions: {option_text}"
reserved_tokens = len(
tokenizer(
prefix + suffix,
add_special_tokens=True,
)["input_ids"]
)
state_ids = tokenizer(
text,
add_special_tokens=False,
)["input_ids"][
:max(1, MAX_LEN - reserved_tokens)
]
prompt = (
prefix
+ tokenizer.decode(
state_ids,
skip_special_tokens=True,
)
+ suffix
)
encoded = tokenizer(
prompt,
max_length=MAX_LEN,
truncation=True,
return_tensors="pt",
)
mask_positions = (
encoded["input_ids"][0]
== tokenizer.mask_token_id
).nonzero(as_tuple=True)[0].tolist()
if len(mask_positions) != len(options):
raise ValueError(
"Option markers were truncated. "
"Use shorter text or fewer options."
)
input_ids = encoded["input_ids"].to(device)
attention_mask = encoded["attention_mask"].to(device)
logits = model(
input_ids,
attention_mask,
[mask_positions],
)[0]
probabilities = F.softmax(
logits.float(),
dim=0,
).cpu().tolist()
scores = dict(zip(options, probabilities))
prediction = max(scores, key=scores.get)
return {
"prediction": prediction,
"confidence": scores[prediction],
"probabilities": scores,
}
def synchronize():
if device == "cuda":
torch.cuda.synchronize()
elif device == "mps":
torch.mps.synchronize()
# ---------------------------------------------------------------------
# Command-line interface
# ---------------------------------------------------------------------
def main():
parser = argparse.ArgumentParser(
description="Run Oryn classification inference."
)
parser.add_argument(
"--text",
default=(
"Your account will be suspended unless "
"you verify your password immediately."
),
)
parser.add_argument(
"--question",
default=(
"Is this message a phishing, scam, "
"or fraud attempt?"
),
)
parser.add_argument(
"--options",
nargs="+",
default=["true", "false"],
)
args = parser.parse_args()
# Warm up the device before measuring latency.
for _ in range(10):
predict(
args.text,
args.question,
args.options,
)
synchronize()
start = time.perf_counter()
result = predict(
args.text,
args.question,
args.options,
)
synchronize()
latency_ms = (
time.perf_counter() - start
) * 1000
print(f"Device: {device}")
print(f"Prediction: {result['prediction']}")
print(f"Confidence: {result['confidence']:.4f}")
print(f"Latency: {latency_ms:.2f} ms")
print("Probabilities:")
for option, probability in sorted(
result["probabilities"].items(),
key=lambda item: item[1],
reverse=True,
):
print(f" {option}: {probability:.4f}")
if __name__ == "__main__":
main()