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

Function main

test/humaneval/humaneval_gen.py:76–172  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

74
75
76def 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

Callers 1

humaneval_gen.pyFile · 0.70

Calls 2

generate_promptFunction · 0.85
get_modelFunction · 0.70

Tested by

no test coverage detected