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

Function main

data/generation/generate_vllm.py:9–32  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

7import json
8
9def main(args):
10 torch.manual_seed(args.seed)
11 world_size = torch.cuda.device_count()
12 n_gpus = torch.cuda.device_count()
13 print(f"using {n_gpus} GPUs to generate")
14
15 model = LLM(model=args.base_model, tensor_parallel_size=n_gpus)
16
17 tokenizer = AutoTokenizer.from_pretrained(args.base_model, use_fast=False)
18
19 prompts, _ = get_gen_dataset(args.dataset_name, args.max_sample, tokenizer)
20
21 sampling_params = SamplingParams(temperature=args.temperature, top_p=1, max_tokens=args.max_new_tokens)
22
23 with torch.no_grad():
24 outputs = model.generate(prompts, sampling_params)
25
26 all_outputs = []
27 for output in outputs:
28 all_outputs.append([[output.prompt, output.outputs[0].text]])
29
30 with open(args.out_path + f'/{args.dataset_name}_T{args.temperature}_N{args.max_new_tokens}_S{args.seed}_{args.max_sample}.json', 'w') as f:
31 for item in all_outputs[:len(outputs)]:
32 f.write(json.dumps(item) + '\n')
33
34if __name__ == "__main__":
35 parser = argparse.ArgumentParser(description='Parameters')

Callers 1

generate_vllm.pyFile · 0.70

Calls 1

get_gen_datasetFunction · 0.90

Tested by

no test coverage detected