MCPcopy Create free account
hub / github.com/OpenBitSys/BitDistiller / get_wikitext2

Function get_wikitext2

test/general/wiki_ppl.py:13–42  ·  view source on GitHub ↗
(nsamples, seed, seqlen, model)

Source from the content-addressed store, hash-verified

11
12
13def get_wikitext2(nsamples, seed, seqlen, model):
14 from datasets import load_dataset
15
16 # traindata = load_dataset('/root/model/datasets/wikitext/wikitext-2-raw-v1', split='train')
17 # testdata = load_dataset('/root/model/datasets/wikitext/wikitext-2-raw-v1', split='test')
18
19 traindata = load_dataset('wikitext', 'wikitext-2-raw-v1', split='train')
20 testdata = load_dataset('wikitext', 'wikitext-2-raw-v1', split='test')
21
22 try:
23 tokenizer = AutoTokenizer.from_pretrained(model, use_fast=False)
24 except:
25 tokenizer = AutoTokenizer.from_pretrained(model, use_fast=True)
26 trainenc = tokenizer("\n\n".join(traindata['text']), return_tensors='pt')
27 testenc = tokenizer("\n\n".join(testdata['text']), return_tensors='pt')
28
29 import random
30 random.seed(seed)
31 np.random.seed(0)
32 torch.random.manual_seed(0)
33
34 traindataset = []
35 for _ in range(nsamples):
36 i = random.randint(0, trainenc.input_ids.shape[1] - seqlen - 1)
37 j = i + seqlen
38 inp = trainenc.input_ids[:, i:j]
39 attention_mask = torch.ones_like(inp)
40 traindataset.append({'input_ids':inp,'attention_mask': attention_mask})
41
42 return traindataset, testenc
43
44@torch.no_grad()
45def llama_eval(model, testenc, dev, seqlen = 2048):

Callers 1

mainFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected