| 20 | |
| 21 | # Load and process aligned dataset |
| 22 | def 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 |