| 76 | return [cat_samples[:, i*block_size:(i+1)*block_size] for i in range(n_split)] |
| 77 | |
| 78 | def get_calib_dataset_gsm8k(tokenizer=None, n_samples=512, block_size=512): |
| 79 | # download from here: https://github.com/OFA-Sys/gsm8k-ScRel/blob/main/data/train_use.jsonl |
| 80 | data_path = "/root/model/gsm8k-ScRel/data/train_use.jsonl" |
| 81 | |
| 82 | with open(data_path, 'r') as f: |
| 83 | dataset_for_eval = f.readlines() |
| 84 | |
| 85 | dataset = [json.loads(item.strip()) for item in dataset_for_eval] |
| 86 | random.seed(42) |
| 87 | dataset = random.sample(dataset, k=min(n_samples * 10, len(dataset))) |
| 88 | samples = [] |
| 89 | n_run = 0 |
| 90 | |
| 91 | for data in dataset: |
| 92 | istr = data["query"] |
| 93 | opt = data["response"] |
| 94 | line = f"Instruction:\n{istr}\nOutput:\n" |
| 95 | line += opt |
| 96 | line = line.strip() |
| 97 | line_encoded = tokenizer.encode(line) |
| 98 | if len(line_encoded) > 512: |
| 99 | continue |
| 100 | sample = torch.tensor([line_encoded]) |
| 101 | if sample.numel() == 0: |
| 102 | continue |
| 103 | samples.append(sample) |
| 104 | n_run += 1 |
| 105 | if n_run == n_samples: |
| 106 | break |
| 107 | # now concatenate all samples and split according to block size |
| 108 | cat_samples = torch.cat(samples, dim=1) |
| 109 | n_split = cat_samples.shape[1] // block_size |
| 110 | print(f" * Split into {n_split} blocks") |
| 111 | |
| 112 | return [cat_samples[:, i*block_size:(i+1)*block_size] for i in range(n_split)] |
| 113 | |
| 114 | def get_blocks(model): |
| 115 | if isinstance(model, LlamaForCausalLM): |