def plot_prob_distribution(probabilities, next_tokens, sampling, potential_nb, total_nb=50): # Get top k tokens top_k_prob, top_k_indices = torch.topk(probabilities, total_nb) top_k_tokens = [tokenizer.decode([idx]) for idx in top_k_indices.tolist()] # Get next tokens and their probabilities next_tokens_list = [tokenizer.decode([idx]) for idx in next_tokens.tolist()] next_token_prob = probabilities[next_tokens].tolist() # Create figure plt.figure(figsize=(0.4*total_nb, 5), dpi=300, facecolor='white') plt.rc('axes', axisbelow=True) plt.grid(axis='y', linestyle='-', alpha=0.5) if potential_nb < total_nb: plt.axvline(x=potential_nb-0.5, ls=':', color='grey', label='Sampled tokens') plt.bar(top_k_tokens, top_k_prob.tolist(), color='blue') plt.bar(next_tokens_list, next_token_prob, color='red', label='Selected tokens') plt.xticks(rotation=45, ha='right', va='top') plt.gca().spines['top'].set_visible(False) plt.gca().spines['right'].set_visible(False) if sampling == 'top_k': plt.title('Probability distribution of predicted tokens with top-k sampling') elif sampling == 'nucleus': plt.title('Probability distribution of predicted tokens with nucleus sampling') plt.legend() plt.savefig(f'{sampling}_{time.time()}.png', dpi=300) plt.close() def top_k_sampling(logits, temperature, top_k, beams, plot=True): assert top_k >= 1 assert beams <= top_k indices_to_remove = logits < torch.topk(logits, top_k)[0][..., -1, None] new_logits = torch.clone(logits) new_logits[indices_to_remove] = float('-inf') # Convert logits to probabilities probabilities = torch.nn.functional.softmax(new_logits / temperature, dim=-1) # Sample n tokens from the resulting distribution 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, 'top_k', top_k) return next_tokens # Start generating text beam_search(input_ids, 0, bar, length, beams, 'top_k', 1)