def nucleus_sampling(logits, temperature, p, beams, plot=True): assert p > 0 assert p <= 1 # Sort the probabilities in descending order and compute cumulative probabilities sorted_logits, sorted_indices = torch.sort(logits, descending=True) probabilities = torch.nn.functional.softmax(sorted_logits / temperature, dim=-1) cumulative_probabilities = torch.cumsum(probabilities, dim=-1) # Create a mask for probabilities that are in the top-p mask = cumulative_probabilities < p # If there's not n index where cumulative_probabilities < p, we use the top n tokens instead if mask.sum() > beams: top_p_index_to_keep = torch.where(mask)[0][-1].detach().cpu().tolist() else: top_p_index_to_keep = beams # Only keep top-p indices indices_to_remove = sorted_indices[top_p_index_to_keep:] sorted_logits[indices_to_remove] = float('-inf') # Sample n tokens from the resulting distribution probabilities = torch.nn.functional.softmax(sorted_logits / temperature, dim=-1) next_tokens = torch.multinomial(probabilities, beams) # Plot distribution if plot: total_prob = torch.nn.functional.softmax(logits / temperature, dim=-1) plot_prob_distribution(total_prob, next_tokens, 'nucleus', top_p_index_to_keep) return next_tokens # Start generating text beam_search(input_ids, 0, bar, length, beams, 'nucleus', 1)