BioJev / example_inference.py
Gabriel382's picture
Release BioJev-4B full checkpoint
d99afc9
Raw History Blame Contribute Delete
1.42 kB
import torch
from peft import PeftConfig, PeftModelForSequenceClassification
from transformers import AutoModelForSequenceClassification, AutoTokenizer
MODEL = "Gabriel382/BioJev"
LABEL2ID = {"contradiction": 0, "entailment": 1, "neutral": 2}
ID2LABEL = {v: k for k, v in LABEL2ID.items()}
cfg = PeftConfig.from_pretrained(MODEL)
base = AutoModelForSequenceClassification.from_pretrained(
cfg.base_model_name_or_path,
num_labels=3,
label2id=LABEL2ID,
id2label=ID2LABEL,
dtype=torch.bfloat16 if torch.cuda.is_available() else torch.float32,
device_map="auto" if torch.cuda.is_available() else None,
)
model = PeftModelForSequenceClassification.from_pretrained(base, MODEL, is_trainable=False)
tokenizer = AutoTokenizer.from_pretrained(MODEL)
if tokenizer.pad_token_id is None:
tokenizer.pad_token = tokenizer.eos_token
model.config.pad_token_id = tokenizer.pad_token_id
model.eval()
premise = "The clinical report describes a bacterial pneumonia."
hypothesis = "The patient has an infectious pulmonary disease."
inputs = tokenizer(premise, hypothesis, return_tensors="pt", truncation=True, max_length=2048)
device = next(model.parameters()).device
inputs = {k: v.to(device) for k, v in inputs.items()}
with torch.inference_mode():
probs = torch.softmax(model(**inputs).logits[0].float(), dim=-1)
for label, idx in LABEL2ID.items():
print(f"{label:14s}: {probs[idx].item():.4f}")