Views
No views yet
1import torch
2from faesm.progen2 import ProGenForCausalLM
3from transformers import AutoTokenizer
4device = 'cuda' if torch.cuda.is_available() else 'cpu'
5model = ProGenForCausalLM.from_pretrained("jinyuan22/ProGen2-small").to(torch.float16).to(device).eval()
6tokenizer = AutoTokenizer.from_pretrained("jinyuan22/ProGen2-small")
7
8# sequence = "1" + "ACDEFGHIKLMNPQRSTVWY" * 50 + "2" # 1002 token
9
10sequence = "2GFLPFRGADEGLAAREAATLAARGTAARAYREDSWAVPVPRGLLGDLTARVAALGAASPPPADPLAVTLDLHHVTAEVALTTVLDAATLVHGQTRVLSAEDAAEAATAAAAATEAYLERLQDFVLFMSASVRVWRRGNAAGATGPEWDQWYTVADRDALGSAPTHLAVLGRQADALCHFVLDRVAWGTCGTPLWSGDEDLGNVVATFAGYADRLATAPRDLIM1"
11
12inputs = tokenizer(sequence, return_tensors="pt").to(device)
13
14with torch.no_grad():
15 logits = model(inputs.input_ids, labels=inputs.input_ids).logits
16
17logits = logits[0][:-1, ...]
18target = inputs.input_ids[0, 1:]
19
20# remove unused logits
21first_token, last_token = 5, 29
22logits = logits[:, first_token:(last_token+1)]
23target = target - first_token
24
25ce_eval = torch.nn.functional.cross_entropy(input=logits.view(-1, logits.size(-1)), target=target.view(-1), reduction="mean").item()
26print(ce_eval)
27assert abs(ce_eval - 2.4) < 0.1