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