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)
| 262 | |
| 263 | |
| 264 | def 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) |
no test coverage detected