MCPcopy Create free account
hub / github.com/OpenBitSys/BitDistiller / get_calib_dataset_gsm8k

Function get_calib_dataset_gsm8k

quantization/clip_utils.py:78–112  ·  view source on GitHub ↗
(tokenizer=None, n_samples=512, block_size=512)

Source from the content-addressed store, hash-verified

76 return [cat_samples[:, i*block_size:(i+1)*block_size] for i in range(n_split)]
77
78def 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
114def get_blocks(model):
115 if isinstance(model, LlamaForCausalLM):

Callers 1

get_calib_datasetFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected