MCPcopy Create free account
hub / github.com/OpenSparseLLMs/Linear-MoE / beam_search

Function beam_search

linear_moe/generation/api.py:264–317  ·  view source on GitHub ↗

Perform beam search to generate sequences. Args: model (torch.nn.Module): The model used for beam search. prompts (List[List[int]], optional): List of prompts, where each prompt is a list of token ids. tokens_to_generate (int, optional): Number of tokens to generate

(model,
                prompts=None,
                tokens_to_generate=0,
                beam_size=0,
                add_BOS=False,
                stop_token=50256,
                num_return_gen=1,
                length_penalty=1,
                prevent_newline_after_colon=False)

Source from the content-addressed store, hash-verified

262
263
264def beam_search(model,
265 prompts=None,
266 tokens_to_generate=0,
267 beam_size=0,
268 add_BOS=False,
269 stop_token=50256,
270 num_return_gen=1,
271 length_penalty=1,
272 prevent_newline_after_colon=False):
273 """
274 Perform beam search to generate sequences.
275
276 Args:
277 model (torch.nn.Module): The model used for beam search.
278 prompts (List[List[int]], optional): List of prompts, where each prompt is a list of token ids.
279 tokens_to_generate (int, optional): Number of tokens to generate.
280 beam_size (int, optional): Beam size for beam search.
281 add_BOS (bool, optional): Whether to add the BOS token to the prompt.
282 stop_token (int, optional): Token that indicates the end of generation.
283 num_return_gen (int, optional): Number of generated sequences to return.
284 length_penalty (float, optional): Length penalty for beam search.
285 prevent_newline_after_colon (bool, optional): Whether to prevent newline.
286
287 Returns:
288 torch.Tensor: The generated tokens.
289 """
290 # Make sure input params are avaialble to all ranks.
291 values = [
292 tokens_to_generate, beam_size, add_BOS, stop_token, num_return_gen,
293 length_penalty, prevent_newline_after_colon
294 ]
295 values_float_tensor = broadcast_float_list(len(values), float_list=values)
296 tokens_to_generate = int(values_float_tensor[0].item())
297 beam_size = int(values_float_tensor[1].item())
298 add_BOS = bool(values_float_tensor[2].item())
299 stop_token = int(values_float_tensor[3].item())
300 num_return_gen = int(values_float_tensor[4].item())
301 length_penalty = values_float_tensor[5].item()
302 prevent_newline_after_colon = values_float_tensor[6].item()
303
304 context_tokens_tensor, context_length_tensor = tokenize_prompts(
305 prompts=prompts,
306 tokens_to_generate=tokens_to_generate,
307 add_BOS=add_BOS)
308
309 return beam_search_and_return_on_first_stage(
310 model,
311 context_tokens_tensor,
312 context_length_tensor,
313 beam_size,
314 stop_token=stop_token,
315 num_return_gen=num_return_gen,
316 length_penalty=length_penalty,
317 prevent_newline_after_colon=prevent_newline_after_colon)

Callers 1

Calls 2

tokenize_promptsFunction · 0.85

Tested by

no test coverage detected