INNER CODE UNIT · Python
LM
HendrikStrobelt/detecting-fake-text · backend/api.py:71
class LM(AbstractLanguageChecker):
def __init__(self, model_name_or_path="gpt2"):
super(LM, self).__init__()
self.enc = GPT2Tokenizer.from_pretrained(model_name_or_path)
self.model = GPT2LMHeadModel.from_pretrained(model_name_or_path)
self.model.to(self.device)
self.model.eval()
self.start_token = self.enc(self.enc.bos_token, return_tensors='pt').data['input_ids'][0]
print(f"Loaded GPT-2 model! on {self.device}")
def check_probabilities(self, in_text, topk=40):
# Process input
token_ids = self.enc(in_text, return_tensors='pt').data['input_ids'][0]
token_ids = torch.concat([self.start_token, token_ids])
# Forward through the model
output = self.model(token_ids.to(self.device))
all_logits = output.logits[:-1].detach().squeeze()
# construct target and pred