import matplotlib.pyplot as plt import networkx as nx import numpy as np import time def get_log_prob(logits, token_id): # Compute the softmax of the logits probabilities = torch.nn.functional.softmax(logits, dim=-1) log_probabilities = torch.log(probabilities) # Get the log probability of the token token_log_probability = log_probabilities[token_id].item() return token_log_probability def greedy_search(input_ids, node, length=5): if length == 0: return input_ids outputs = model(input_ids) predictions = outputs.logits # Get the predicted next sub-word (here we use top-k search) logits = predictions[0, -1, :] token_id = torch.argmax(logits).unsqueeze(0) # Compute the score of the predicted token token_score = get_log_prob(logits, token_id) # Add the predicted token to the list of input ids new_input_ids = torch.cat([input_ids, token_id.unsqueeze(0)], dim=-1) # Add node and edge to graph next_token = tokenizer.decode(token_id, skip_special_tokens=True) current_node = list(graph.successors(node))[0] graph.nodes[current_node]['tokenscore'] = np.exp(token_score) * 100 graph.nodes[current_node]['token'] = next_token + f"_{length}" # Recursive call input_ids = greedy_search(new_input_ids, current_node, length-1) return input_ids # Parameters length = 5 beams = 1 # Create a balanced tree with height 'length' graph = nx.balanced_tree(1, length, create_using=nx.DiGraph()) # Add 'tokenscore', 'cumscore', and 'token' attributes to each node for node in graph.nodes: graph.nodes[node]['tokenscore'] = 100 graph.nodes[node]['token'] = text # Start generating text output_ids = greedy_search(input_ids, 0, length=length) output = tokenizer.decode(output_ids.squeeze().tolist(), skip_special_tokens=True) print(f"Generated text: {output}")