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

Function main

data/generation/single_generate.py:149–218  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

147 return gen_dataset, data_collator
148
149def main(args):
150 torch.manual_seed(args.seed)
151
152 base_model = args.base_model
153 batch_size = args.batch_size
154 return_seq_num = 1
155
156 model = AutoModelForCausalLM.from_pretrained(
157 base_model,
158 torch_dtype=torch.bfloat16,
159 device_map='auto'
160 )
161
162 tokenizer = AutoTokenizer.from_pretrained(base_model, use_fast=False)
163 if tokenizer.pad_token is None:
164 smart_tokenizer_and_embedding_resize(
165 special_tokens_dict=dict(pad_token=DEFAULT_PAD_TOKEN),
166 tokenizer=tokenizer,
167 model=model,
168 )
169
170 model.eval()
171
172 # Get the generation dataset
173 gen_dataset, data_collator = make_supervised_data_module(tokenizer, args.dataset_name, args.max_sample)
174
175 dataloader = DataLoader(
176 gen_dataset,
177 shuffle=False,
178 collate_fn=data_collator,
179 batch_size=batch_size,
180 drop_last=True
181 )
182
183 generation_config = GenerationConfig(
184 temperature=args.temperature,
185 do_sample=True,
186 num_beams=return_seq_num,
187 max_new_tokens=args.max_new_tokens,
188 num_return_sequences=return_seq_num,
189 top_p=1.0
190 )
191
192 all_outputs = []
193 total_nums = len(gen_dataset) / args.batch_size
194 for step, batch in tqdm(enumerate(dataloader), total=total_nums):
195 input_ids = batch['input_ids'].to(model.device)
196 attention_mask = batch['attention_mask'].to(model.device)
197 with torch.no_grad():
198 generation_output = model.generate(
199 input_ids=input_ids,
200 attention_mask=attention_mask,
201 generation_config=generation_config,
202 return_dict_in_generate=True
203 )
204
205 s = generation_output.sequences
206

Callers 1

single_generate.pyFile · 0.70

Tested by

no test coverage detected