MCPcopy Create free account
hub / github.com/bytedance/ABQ-LLM / get_wikitext2

Function get_wikitext2

algorithm/datautils.py:30–49  ·  view source on GitHub ↗
(nsamples, seed, seqlen, model)

Source from the content-addressed store, hash-verified

28
29
30def get_wikitext2(nsamples, seed, seqlen, model):
31 print("get_wikitext2")
32 traindata = load_dataset('wikitext', 'wikitext-2-raw-v1', split='train')
33 testdata = load_dataset('wikitext', 'wikitext-2-raw-v1', split='test')
34
35 tokenizer = AutoTokenizer.from_pretrained(model, use_fast=False)
36 trainenc = tokenizer("\n\n".join(traindata['text']), return_tensors='pt')
37 testenc = tokenizer("\n\n".join(testdata['text']), return_tensors='pt')
38
39
40 random.seed(seed)
41 trainloader = []
42 for _ in range(nsamples):
43 i = random.randint(0, trainenc.input_ids.shape[1] - seqlen - 1)
44 j = i + seqlen
45 inp = trainenc.input_ids[:, i:j]
46 tar = inp.clone()
47 tar[:, :-1] = -100
48 trainloader.append((inp, tar))
49 return trainloader, testenc
50
51def get_ptb(nsamples, seed, seqlen, model):
52 print("get_ptb")

Callers 1

get_loadersFunction · 0.85

Calls 3

from_pretrainedMethod · 0.45
cloneMethod · 0.45
appendMethod · 0.45

Tested by

no test coverage detected