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