inference/generate.py
5.7 KB · 145 lines · python Raw
1 import os
2 import json
3 import sys
4 from argparse import ArgumentParser
5 from typing import List
6
7 import torch
8 import torch.distributed as dist
9 from transformers import AutoTokenizer
10 from safetensors.torch import load_model
11
12 from model import Transformer, ModelArgs
13 current_dir = os.path.dirname(os.path.abspath(__file__))
14 encoding_dir = os.path.join(current_dir, '../encoding')
15 sys.path.insert(0, os.path.abspath(encoding_dir))
16 from encoding_dsv4 import encode_messages, parse_message_from_completion_text
17
18
19 @torch.inference_mode()
20 def generate(
21 model: Transformer,
22 prompt_tokens: List[List[int]],
23 max_new_tokens: int,
24 eos_id: int,
25 ) -> List[List[int]]:
26 """Batch generation with left-padded prompts.
27
28 The first forward pass processes [min_prompt_len:] tokens (prefill phase).
29 Subsequent passes generate one token at a time (decode phase). For positions
30 still within a prompt, the ground-truth token overrides the model's prediction.
31 """
32 prompt_lens = [len(t) for t in prompt_tokens]
33 assert max(prompt_lens) <= model.max_seq_len, f"Prompt length exceeds model maximum sequence length (max_seq_len={model.max_seq_len})"
34 total_len = min(model.max_seq_len, max_new_tokens + max(prompt_lens))
35 tokens = torch.full((len(prompt_tokens), total_len), -1, dtype=torch.long)
36 for i, t in enumerate(prompt_tokens):
37 tokens[i, :len(t)] = torch.tensor(t, dtype=torch.long)
38 prev_pos = 0
39 finished = torch.tensor([False] * len(prompt_tokens))
40 prompt_mask = tokens != -1
41 for cur_pos in range(min(prompt_lens), total_len):
42 next_token = model.forward(tokens[:, prev_pos:cur_pos], prev_pos)[0]
43 next_token = torch.where(prompt_mask[:, cur_pos], tokens[:, cur_pos], next_token)
44 tokens[:, cur_pos] = next_token
45 finished |= torch.logical_and(~prompt_mask[:, cur_pos], next_token == eos_id)
46 prev_pos = cur_pos
47 if finished.all():
48 break
49 completion_tokens = []
50 for i, toks in enumerate(tokens.tolist()):
51 toks = toks[prompt_lens[i]:prompt_lens[i]+max_new_tokens]
52 if eos_id in toks:
53 toks = toks[:toks.index(eos_id)]
54 toks.append(eos_id)
55 completion_tokens.append(toks)
56 return completion_tokens
57
58
59 def main(
60 ckpt_path: str,
61 config: str,
62 input_file: str = "",
63 interactive: bool = True,
64 max_new_tokens: int = 100,
65 temperature: float = 1.0,
66 ) -> None:
67 world_size = int(os.getenv("WORLD_SIZE", "1"))
68 rank = int(os.getenv("RANK", "0"))
69 local_rank = int(os.getenv("LOCAL_RANK", "0"))
70 if world_size > 1:
71 dist.init_process_group("nccl")
72 global print
73 if rank != 0:
74 print = lambda *_, **__: None
75 torch.cuda.set_device(local_rank)
76 torch.cuda.memory._set_allocator_settings("expandable_segments:True")
77 torch.set_default_dtype(torch.bfloat16)
78 torch.set_num_threads(8)
79 torch.manual_seed(33377335)
80 with open(config) as f:
81 args = ModelArgs(**json.load(f))
82 args.temperature = temperature
83 if interactive:
84 args.max_batch_size = 1
85 args.max_seq_len = 64 * 1024
86 print(args)
87 with torch.device("cuda"):
88 model = Transformer(args)
89 tokenizer = AutoTokenizer.from_pretrained(ckpt_path)
90 print("load model")
91 load_model(model, os.path.join(ckpt_path, f"model{rank}-mp{world_size}.safetensors"), strict=False)
92 torch.set_default_device("cuda")
93 print("I'm DeepSeek 👋")
94
95 if interactive:
96 messages = []
97 while True:
98 if world_size == 1:
99 prompt = input(">>> ")
100 elif rank == 0:
101 prompt = input(">>> ")
102 objects = [prompt]
103 dist.broadcast_object_list(objects, 0)
104 else:
105 objects = [None]
106 dist.broadcast_object_list(objects, 0)
107 prompt = objects[0]
108 if prompt == "/exit":
109 break
110 elif prompt == "/clear":
111 messages.clear()
112 continue
113 messages.append({"role": "user", "content": prompt})
114 prompt_tokens = tokenizer.encode(encode_messages(messages, thinking_mode="chat"))
115 completion_tokens = generate(model, [prompt_tokens], max_new_tokens, tokenizer.eos_token_id)
116 completion = tokenizer.decode(completion_tokens[0])
117 print(completion)
118 messages.append(parse_message_from_completion_text(completion, thinking_mode="chat"))
119 else:
120 with open(input_file) as f:
121 prompts = f.read().split("\n\n")
122 prompt_tokens = [tokenizer.encode(encode_messages([{"role": "user", "content": prompt}], thinking_mode="chat")) for prompt in prompts]
123 completion_tokens = generate(model, prompt_tokens, max_new_tokens, tokenizer.eos_token_id)
124 completions = tokenizer.batch_decode(completion_tokens)
125 for prompt, completion in zip(prompts, completions):
126 print("Prompt:", prompt)
127 print("Completion:", completion)
128 print()
129
130 if world_size > 1:
131 dist.destroy_process_group()
132
133
134 if __name__ == "__main__":
135 parser = ArgumentParser()
136 parser.add_argument("--ckpt-path", type=str, required=True)
137 parser.add_argument("--config", type=str, required=True)
138 parser.add_argument("--input-file", type=str, default="")
139 parser.add_argument("--interactive", action="store_true")
140 parser.add_argument("--max-new-tokens", type=int, default=300)
141 parser.add_argument("--temperature", type=float, default=1.0)
142 args = parser.parse_args()
143 assert args.input_file or args.interactive, "Either input-file or interactive mode must be specified"
144 main(args.ckpt_path, args.config, args.input_file, args.interactive, args.max_new_tokens, args.temperature)
145