| 69 | |
| 70 | |
| 71 | def get_alpaca(nsamples, seed, seqlen, tokenizer, disentangle=False, dataset="alpaca"): |
| 72 | if dataset == "alpaca": |
| 73 | data_files = {"train": "./data/alpaca_train.csv"} |
| 74 | elif dataset == "alpaca_cleaned": |
| 75 | data_files = {"train": "./data/alpaca_cleaned_train.csv"} |
| 76 | elif dataset == "alpaca_cleaned_no_safety": |
| 77 | data_files = {"train": "./data/alpaca_cleaned_no_safety_train.csv"} |
| 78 | else: |
| 79 | raise ValueError("Dataset not supported") |
| 80 | traindata = load_dataset("csv", data_files=data_files, split="train") |
| 81 | random.seed(seed) |
| 82 | # Encode datasets |
| 83 | trainloader = [] |
| 84 | if disentangle: |
| 85 | traindata_sampled = traindata.shuffle(seed=seed).select(range(nsamples)) |
| 86 | for i in range(nsamples): |
| 87 | trainenc_prompt = tokenizer( |
| 88 | traindata_sampled["prompt"][i], return_tensors="pt" |
| 89 | ) |
| 90 | trainenc_response = tokenizer( |
| 91 | traindata_sampled["response"][i], return_tensors="pt" |
| 92 | ) |
| 93 | inp = torch.cat( |
| 94 | (trainenc_prompt.input_ids, trainenc_response.input_ids[:, 1:]), dim=1 |
| 95 | ) # to remove the first token of the response ('1') |
| 96 | tar = inp.clone() |
| 97 | trainenc_prompt_len = trainenc_prompt.input_ids.shape[1] |
| 98 | tar[:, :trainenc_prompt_len] = -100 |
| 99 | trainloader.append((inp, tar)) |
| 100 | else: |
| 101 | trainenc = tokenizer(" ".join(traindata["text"]), return_tensors="pt") |
| 102 | # Generate samples from training set |
| 103 | for _ in range(nsamples): |
| 104 | i = random.randint(0, trainenc.input_ids.shape[1] - seqlen - 1) |
| 105 | j = i + seqlen |
| 106 | inp = trainenc.input_ids[:, i:j] |
| 107 | tar = inp.clone() |
| 108 | tar[:, :-1] = -100 |
| 109 | trainloader.append((inp, tar)) |
| 110 | return trainloader, None |
| 111 | |
| 112 | |
| 113 | # Function to select the appropriate loader based on dataset name |