from tqdm.notebook import tqdm def greedy_sampling(logits, beams): return torch.topk(logits, beams).indices def beam_search(input_ids, node, bar, length, beams, sampling, temperature=0.1): if length == 0: return None outputs = model(input_ids) predictions = outputs.logits # Get the predicted next sub-word (here we use top-k search) logits = predictions[0, -1, :] if sampling == 'greedy': top_token_ids = greedy_sampling(logits, beams) elif sampling == 'top_k': top_token_ids = top_k_sampling(logits, temperature, 20, beams) elif sampling == 'nucleus': top_token_ids = nucleus_sampling(logits, temperature, 0.5, beams) for j, token_id in enumerate(top_token_ids): bar.update(1) # Compute the score of the predicted token token_score = get_log_prob(logits, token_id) cumulative_score = graph.nodes[node]['cumscore'] + token_score # Add the predicted token to the list of input ids new_input_ids = torch.cat([input_ids, token_id.unsqueeze(0).unsqueeze(0)], dim=-1) # Add node and edge to graph token = tokenizer.decode(token_id, skip_special_tokens=True) current_node = list(graph.successors(node))[j] graph.nodes[current_node]['tokenscore'] = np.exp(token_score) * 100 graph.nodes[current_node]['cumscore'] = cumulative_score graph.nodes[current_node]['sequencescore'] = 1/(len(new_input_ids.squeeze())) * cumulative_score graph.nodes[current_node]['token'] = token + f"_{length}_{j}" # Recursive call beam_search(new_input_ids, current_node, bar, length-1, beams, sampling, 1) # Parameters length = 5 beams = 2 # Create a balanced tree with height 'length' and branching factor 'k' graph = nx.balanced_tree(beams, length, create_using=nx.DiGraph()) bar = tqdm(total=len(graph.nodes)) # Add 'tokenscore', 'cumscore', and 'token' attributes to each node for node in graph.nodes: graph.nodes[node]['tokenscore'] = 100 graph.nodes[node]['cumscore'] = 0 graph.nodes[node]['sequencescore'] = 0 graph.nodes[node]['token'] = text # Start generating text beam_search(input_ids, 0, bar, length, beams, 'greedy', 1)