MCPcopy Create free account
hub / github.com/csuhan/OneLLM / stream_generate

Method stream_generate

model/meta.py:115–162  ·  view source on GitHub ↗
(
        self,
        prompt: str,
        images,
        max_gen_len: int,
        temperature: float = 0.8,
        top_p: float = 0.95,
        modal = ['image'],
    )

Source from the content-addressed store, hash-verified

113
114 @torch.inference_mode()
115 def stream_generate(
116 self,
117 prompt: str,
118 images,
119 max_gen_len: int,
120 temperature: float = 0.8,
121 top_p: float = 0.95,
122 modal = ['image'],
123 ):
124 params = self.llma.params
125
126 prompt_tokens = self.tokenizer.encode(prompt, bos=True, eos=False)
127 # truncate from the left. leave some space for generation.
128 max_seq_len = params.max_seq_len
129 if images is not None:
130 max_seq_len -= self.llma.image_words
131
132 max_prompt_size = max_seq_len - max_gen_len
133 prompt_tokens = prompt_tokens[-max_prompt_size:]
134
135 prompt_size = len(prompt_tokens)
136
137 total_len = min(max_seq_len, max_gen_len + prompt_size)
138
139 tokens = torch.full([total_len], 0).cuda().long()
140
141 tokens[:len(prompt_tokens)] = torch.tensor(prompt_tokens).long()
142 start_pos = prompt_size
143 prev_pos = 0
144 generate_until = start_pos
145 for cur_pos in range(start_pos, total_len):
146 logits = self.llma.forward_inference(tokens[None, prev_pos:cur_pos], prev_pos, images if prev_pos == 0 else None, modal = modal)
147 if temperature > 0:
148 probs = torch.softmax(logits / temperature, dim=-1)
149 next_token = self.sample_top_p(probs, top_p)
150 else:
151 next_token = torch.argmax(logits, dim=-1)
152 next_token = next_token.item()
153
154 if next_token == self.tokenizer.eos_id:
155 break
156
157 tokens[cur_pos] = next_token
158 prev_pos = cur_pos
159 generate_until = cur_pos + 1
160 yield {"text": self.tokenizer.decode(tokens[start_pos:generate_until].tolist()), "end_of_content": False}
161
162 yield {"text": self.tokenizer.decode(tokens[start_pos:generate_until].tolist()), "end_of_content": True}
163
164 def sample_top_p(self, probs, p):
165 probs_sort, probs_idx = torch.sort(probs, dim=-1, descending=True)

Callers 1

model_workerFunction · 0.95

Calls 4

sample_top_pMethod · 0.95
encodeMethod · 0.80
forward_inferenceMethod · 0.80
decodeMethod · 0.80

Tested by

no test coverage detected