()
| 74 | |
| 75 | |
| 76 | def main(): |
| 77 | parser = argparse.ArgumentParser() |
| 78 | |
| 79 | parser.add_argument('--model', type=str, default='bigcode/starcoder', help="") |
| 80 | parser.add_argument('--output_path', type=str, help="") |
| 81 | parser.add_argument('--start_index', type=int, default=0, help="") |
| 82 | parser.add_argument('--end_index', type=int, default=164, help="") |
| 83 | parser.add_argument('--temperature', type=float, default=0.8, help="") |
| 84 | parser.add_argument('--N', type=int, default=200, help="") |
| 85 | parser.add_argument('--max_len', type=int, default=512, help="") |
| 86 | parser.add_argument('--decoding_style', type=str, default='sampling', help="") |
| 87 | parser.add_argument('--num_seqs_per_iter', type=int, default=50, help='') |
| 88 | parser.add_argument('--greedy_decode', action='store_true', help='') |
| 89 | parser.add_argument('--overwrite', default=True, help='') |
| 90 | parser.add_argument("--seed", type=int, default=42, help="seed") |
| 91 | parser.add_argument('--quant_type', type=str, default='n2f3', help='quantization type') |
| 92 | parser.add_argument('--bits', type=int, default=4, help="") |
| 93 | parser.add_argument('--group_size', type=int, default=128, help="") |
| 94 | args = parser.parse_args() |
| 95 | torch.manual_seed(args.seed) |
| 96 | argsdict = vars(args) |
| 97 | print(pprint.pformat(argsdict)) |
| 98 | |
| 99 | problems = read_problems() |
| 100 | |
| 101 | task_ids = sorted(problems.keys())[args.start_index: args.end_index] |
| 102 | prompts = [problems[task_id]['prompt'] for task_id in task_ids] |
| 103 | num_samples = len(prompts) |
| 104 | print("Number of samples: {}".format(num_samples)) |
| 105 | |
| 106 | tokenizer, model = get_model(base_model=args.model, quant_type=args.quant_type, group_size=args.group_size, bits=args.bits, args=args) |
| 107 | generation_config = GenerationConfig( |
| 108 | pad_token_id=tokenizer.pad_token_id, |
| 109 | do_sample=False if args.greedy_decode else True, |
| 110 | temperature=args.temperature, |
| 111 | max_length=args.max_len, |
| 112 | num_return_sequences=args.num_seqs_per_iter, |
| 113 | eos_token_id=tokenizer.eos_token_id, |
| 114 | top_p=0.95 |
| 115 | ) |
| 116 | |
| 117 | print(f"Loaded {args.model}.") |
| 118 | for i in tqdm(range(num_samples), ncols=0, total=num_samples): |
| 119 | output_file = args.output_path + '/{}.jsonl'.format(args.start_index + i) |
| 120 | |
| 121 | if os.path.exists(output_file) and not args.overwrite: |
| 122 | print(f'Skip {output_file} as it already exists') |
| 123 | continue |
| 124 | |
| 125 | prompt = prompts[i].replace(' ', '\t') |
| 126 | prompt_batch = [generate_prompt(prompt)] |
| 127 | |
| 128 | ids_batch = [task_ids[i]] |
| 129 | |
| 130 | completion_seqs = [] |
| 131 | |
| 132 | encoding = tokenizer(prompt_batch, return_tensors="pt", truncation=True, max_length=args.max_len).to(device) |
| 133 |
no test coverage detected