()
| 83 | |
| 84 | |
| 85 | def main(): |
| 86 | parser = argparse.ArgumentParser() |
| 87 | parser = add_code_generation_args(parser) |
| 88 | args, _ = parser.parse_known_args() |
| 89 | |
| 90 | print("Loading tokenizer ...") |
| 91 | tokenizer = CodeGeeXTokenizer( |
| 92 | tokenizer_path=args.tokenizer_path, |
| 93 | mode="codegeex-13b") |
| 94 | |
| 95 | print("Loading state dict ...") |
| 96 | state_dict = torch.load(args.load, map_location="cpu") |
| 97 | state_dict = state_dict["module"] |
| 98 | |
| 99 | print("Building CodeGeeX model ...") |
| 100 | model = model_provider(args) |
| 101 | model.load_state_dict(state_dict) |
| 102 | model.eval() |
| 103 | model.half() |
| 104 | if args.quantize: |
| 105 | model = quantize(model, weight_bit_width=8, backend="torch") |
| 106 | model.cuda() |
| 107 | |
| 108 | def predict( |
| 109 | prompt, |
| 110 | lang, |
| 111 | seed, |
| 112 | out_seq_length, |
| 113 | temperature, |
| 114 | top_k, |
| 115 | top_p, |
| 116 | ): |
| 117 | set_random_seed(seed) |
| 118 | if lang.lower() in LANGUAGE_TAG: |
| 119 | prompt = LANGUAGE_TAG[lang.lower()] + "\n" + prompt |
| 120 | |
| 121 | generated_code = codegeex.generate( |
| 122 | model, |
| 123 | tokenizer, |
| 124 | prompt, |
| 125 | out_seq_length=out_seq_length, |
| 126 | seq_length=args.max_position_embeddings, |
| 127 | top_k=top_k, |
| 128 | top_p=top_p, |
| 129 | temperature=temperature, |
| 130 | micro_batch_size=args.micro_batch_size, |
| 131 | backend="megatron", |
| 132 | verbose=True, |
| 133 | ) |
| 134 | return prompt + generated_code |
| 135 | |
| 136 | examples = [] |
| 137 | with open(args.example_path, "r") as f: |
| 138 | for line in f: |
| 139 | examples.append(list(json.loads(line).values())) |
| 140 | |
| 141 | with gr.Blocks() as demo: |
| 142 | gr.Markdown( |
no test coverage detected