File size: 1,587 Bytes
02030c8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
"""Minimal BioJev-Nano inference example."""

import torch
from peft import PeftConfig, PeftModelForSequenceClassification
from transformers import AutoModelForSequenceClassification, AutoTokenizer

MODEL = "Gabriel382/BioJev-Nano"

LABEL2ID = {
    "contradiction": 0,
    "entailment": 1,
    "neutral": 2,
}
ID2LABEL = {v: k for k, v in LABEL2ID.items()}

peft_config = PeftConfig.from_pretrained(MODEL)

base = AutoModelForSequenceClassification.from_pretrained(
    peft_config.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=512,
)

device = next(model.parameters()).device
inputs = {k: v.to(device) for k, v in inputs.items()}

with torch.inference_mode():
    logits = model(**inputs).logits[0].float()
    probs = torch.softmax(logits, dim=-1)

print("BioJev-Nano")
for label, idx in LABEL2ID.items():
    print(f"{label:14s}: {probs[idx].item():.4f}")