MCPcopy Create free account
hub / github.com/boyiwei/alignment-attribution-code / get_align

Function get_align

lib/data.py:22–59  ·  view source on GitHub ↗
(nsamples, seed, seqlen, tokenizer, disentangle=False, mode="base")

Source from the content-addressed store, hash-verified

20
21# Load and process aligned dataset
22def get_align(nsamples, seed, seqlen, tokenizer, disentangle=False, mode="base"):
23 # Load train and test datasets
24 if mode == "short":
25 data_files = {"train": "./data/SFT_aligned_llama2-7b-chat-hf_train_short.csv"}
26 else:
27 data_files = {"train": "./data/SFT_aligned_llama2-7b-chat-hf_train.csv"}
28 traindata = load_dataset("csv", data_files=data_files, split="train")
29 trainloader = []
30 random.seed(seed)
31 if disentangle:
32 traindata_sampled = traindata.shuffle(seed=seed).select(range(nsamples))
33 for i in range(nsamples):
34 trainenc_prompt = tokenizer(
35 traindata_sampled["prompt"][i], return_tensors="pt"
36 )
37 trainenc_response = tokenizer(
38 traindata_sampled["response"][i], return_tensors="pt"
39 )
40 inp = torch.cat(
41 (trainenc_prompt.input_ids, trainenc_response.input_ids[:, 1:]), dim=1
42 )
43 tar = inp.clone()
44 trainenc_prompt_len = trainenc_prompt.input_ids.shape[1]
45 tar[:, :trainenc_prompt_len] = -100
46 trainloader.append((inp, tar))
47 else:
48 # Encode datasets
49 trainenc = tokenizer(" ".join(traindata["text"]), return_tensors="pt")
50
51 # Generate samples from training set
52 for _ in range(nsamples):
53 i = random.randint(0, trainenc.input_ids.shape[1] - seqlen - 1)
54 j = i + seqlen
55 inp = trainenc.input_ids[:, i:j]
56 tar = inp.clone()
57 tar[:, :-1] = -100
58 trainloader.append((inp, tar))
59 return trainloader, None
60
61
62# Load and process wikitext2 dataset

Callers 1

get_loadersFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected