1import torch
2from src.bert_classifier import BertClassifier
3
4# Load checkpoint
5ckpt = torch.load('pytorch_model.bin', map_location='cpu')
6model = BertClassifier(
7 vocab_size=5175,
8 encoder_num=6,
9)
10model.load_state_dict(ckpt['model_state_dict'])
11model.eval()
12
13# Inference
14input_ids = tokenizer(text, return_tensors='pt')['input_ids']
15logits = model(input_ids)
16pred = torch.argmax(logits, dim=-1).item()