diff --git a/.gitignore b/.gitignore index 293af7e..3687207 100644 --- a/.gitignore +++ b/.gitignore @@ -129,6 +129,8 @@ dmypy.json .pyre/ # Other +*test.py +*test.ipynb *alignment_texts.jsonl *config.py *.DS_Store diff --git a/src/dataset/create_dataset.py b/src/dataset/create_dataset.py index 2456e44..8834c07 100644 --- a/src/dataset/create_dataset.py +++ b/src/dataset/create_dataset.py @@ -1,4 +1,3 @@ -import jsonlines import numpy as np from typing import List, Dict, Tuple, DefaultDict, Any from collections import defaultdict @@ -11,6 +10,9 @@ from pathlib import Path from tqdm.auto import tqdm from dateutil.parser import parse, ParserError import openai +from datasets import load_dataset +from langchain.document_loaders import HuggingFaceDatasetLoader + try: import config @@ -19,18 +21,18 @@ except ImportError: openai.api_key = os.environ.get('OPENAI_API_KEY') -from .settings import PATH_TO_RAW_DATA, PATH_TO_DATASET_PKL, PATH_TO_DATASET_DICT_PKL, EMBEDDING_MODEL, LEN_EMBEDDINGS +from .settings import PATH_TO_DATASET_PKL, PATH_TO_DATASET_DICT_PKL, EMBEDDING_MODEL, LEN_EMBEDDINGS from .text_splitter import TokenSplitter, split_into_sentences error_count_dict = { + "Entry has no id.": 0, "Entry has no source.": 0, "Entry has no title.": 0, "Entry has no text.": 0, - "Entry has no URL.": 0, - "Entry has wrong citation level.": 0 + "Entry has non-string required key.": 0 } @@ -38,221 +40,129 @@ class MissingDataException(Exception): pass -class Dataset: +class ChunkedARD: def __init__(self, - jsonl_data_path: str = PATH_TO_RAW_DATA, # Path to the dataset .jsonl file. - custom_sources: List[str] = None, # List of sources to include, like "alignment forum", "lesswrong", "arxiv",etc. - rate_limit_per_minute: int = 3_500, # Rate limit for the OpenAI API. min_tokens_per_block: int = 300, # Minimum number of tokens per block. max_tokens_per_block: int = 400, # Maximum number of tokens per block. - fraction_of_articles_to_use: float = 1.0, # Fraction of articles to use. If 1.0, use all articles. - ): - self.jsonl_data_path = jsonl_data_path + custom_sources: List[str] = None, # List of sources to include, like "alignmentforum", "lesswrong", "arxiv",etc. + rate_limit_per_minute: int = 3_500, # Rate limit for the OpenAI API. + ): + self.min_tokens_per_block = min_tokens_per_block # for the text splitter + self.max_tokens_per_block = max_tokens_per_block # for the text splitter + self.custom_sources = custom_sources self.rate_limit_per_minute = rate_limit_per_minute self.delay_in_seconds = 60.0 / self.rate_limit_per_minute - self.fraction_of_articles_to_use = fraction_of_articles_to_use - self.min_tokens_per_block = min_tokens_per_block # for the text splitter - self.max_tokens_per_block = max_tokens_per_block # for the text splitter - - self.metadata: List[Tuple[str]] = [] # List of tuples, each containing the title, author, date, URL, and tags of an article. - self.embedding_strings: List[str] = [] # List of strings, each being a few paragraphs from a single article (not exceeding max_tokens_per_block tokens). - self.embeddings_metadata_index: List[int] = [] # List of integers, each being the index of the article from which the embedding string was taken. + self.metadata: List[Dict[str, Any]] = [] # List of dicts, each containing: id, entry_id, source, title, text, url, date_published, authors. - self.articles_count: DefaultDict[str, int] = defaultdict(int) # Number of articles per source. E.g.: {'source1': 10, 'source2': 20, 'total': 30} + self.entries_per_source_count: DefaultDict[str, int] = defaultdict(int) # Number of entries per source. E.g.: {'source1': 10, 'source2': 20, 'total': 30} + self.total_counts = { + 'chars': 0, + 'words': 0, + 'sentences': 0, + 'chunks': 0, + 'entries': 0, + } if self.custom_sources is not None: for source in self.custom_sources: - self.articles_count[source] = 0 - self.total_articles_count = 0 + self.entries_per_source_count[source] = 0 - self.total_char_count = 0 - self.total_word_count = 0 - self.total_sentence_count = 0 - self.total_block_count = 0 - - self.sources_so_far: List[str] = [] - self.info_types: Dict[str, List[str]] = {} - - def extract_info_from_article(self, article: Dict[str, Any]) -> Tuple[str]: - """ - This function extracts the title, author, date, URL, tags, and text from an article. - - Args: - article (Dict[str, Any]): a dictionary containing the article's text and metadata. + def contains_required_metadata(self, entry: Dict[str, Any]): + metadata_types = { + 'id': str, + 'source': str, + 'title': str, + 'url': str, + 'date_published': str, + 'authors': list, # It appears to be a list, but actually it's a list-looking string. Like "['apple', 'orange', 'tomato']" is a string. +# 'summary': str # see previous comment + } + required_metadata_keys = ['id', 'source', 'title', 'text'] - Returns: - Tuple[str]: a tuple containing the title, author, date, URL, tags, and text of the article. - """ - title: str = "" - author: str = "" - date_published: str = None - url: str = None - tags: str = None - text: str = None - - # Get title - if 'title' in article and 'book_title' in article and article['title']: title = article['title'] - elif 'book_title' in article and 'title' not in article and article['book_title']: - title = article['book_title'] - elif 'title' in article and article['title']: - title = article['title'] - title = title.strip('\n').replace('\n', ' ')[:100] + # Check that the 8 primary metadata keys all have the correct type + for key, key_type in metadata_types.items(): + if type(entry[key]) != key_type: + raise MissingDataException(f"Entry {entry['id']} has key {key} of type {type(entry[key])} when it should be {key_type}.") - # Get author - if 'author' in article and 'authors' in article and article['author']: author = article['author'] - elif 'authors' in article and article['authors']: author = article['authors'] - elif 'author' in article and article['author']: author = article['author'] - if type(author) == str: author = get_authors_list(author) - if type(author) == list: author = ', '.join(author) - author = author.strip('\n').replace('\n', ' ')[:100] + # Check that the 4 required metadata keys are non-empty + for key in required_metadata_keys: + if not entry[key]: + raise MissingDataException(f"Entry {entry['id']} has no {key}.") - # Get date published - if 'date_published' in article and article['date_published'] and len(article['date_published']) >= 10: date_published = article['date_published'][:10] - elif 'published' in article and article['published'] and len(article['published']) >= 16: date_published = article['published'][:16] - else: date_published = None - if date_published is not None: - date_published = standardize_date(date_published) - - # Get URL - if 'link' in article and article['link']: url = article['link'] - elif 'url' in article and article['url']: url = article['url'] - elif 'doi' in article and article['doi']: url = article['doi'] - else: url = None - - # Get tags - if 'tags' in article and article['tags']: - if type(article['tags']) == list: tags = ', '.join([val['term'] for val in article['tags']]) - elif type(article['tags']) == str: tags = article['tags'] - else: tags = None - - # Get text - if 'text' in article and article['text']: text = article['text'] - else: - raise MissingDataException(f"Entry has no text.") - - return (title, author, date_published, url, tags, text) def get_alignment_texts(self): text_splitter = TokenSplitter(self.min_tokens_per_block, self.max_tokens_per_block) - with jsonlines.open(self.jsonl_data_path, "r") as reader: - for entry in tqdm(reader): - try: - if 'source' not in entry: - if 'url' in entry and entry['url'] == "https://www.cold-takes.com/": - entry["source"] = "Cold Takes" - elif 'question' in entry and 'answer' in entry: - entry["source"] = "printouts" - continue # for now, skip printouts - elif 'article_url' in entry and entry['article_url'] == "https://www.gwern.net": - entry["source"] = "gwern.net" - elif 'url' in entry and entry['url'] == "https://generative.ink/posts/": - entry["source"] = "generative.ink" - elif 'url' in entry and entry['url'][:24] == "https://greaterwrong.com": - entry["source"] = "greaterwrong.com" - else: - raise MissingDataException("Entry has no source.") - - # if we specified custom sources, only include articles from those sources - if (self.custom_sources is not None) and (entry['source'] not in self.custom_sources): - continue - - if entry["source"] == 'alignment forum': - if int(entry['score'].replace('−', '-')) < 70: continue - elif entry["source"] == 'lesswrong': - if int(entry['score'].replace('−', '-')) < 150: continue - elif entry["source"] == 'arxiv': - if 'citation_level' != '0': continue + # Load the dataset. streaming allows you to load one entry at a time, + # so entries can be processed before the entire dataset has been saved. + iterable_data = load_dataset('StampyAI/alignment-research-dataset', 'aisafety.info', split='train', streaming=True) + + for entry in tqdm(iterable_data): + """Checks""" - # Dict describing the proportion of each source we want: - # E.g.: {'arxiv': 0.5, 'youtube': 0.5, 'lesswrong': 1.0} - desired_source_proportions = { - "https://aipulse.org": 1, - "ebook": 0, - "https://qualiacomputing.com": 0.02, - "alignment forum": .7, - "lesswrong": .5, - "manual": 1, - "arxiv": 0.1, - "https://deepmindsafetyresearch.medium.com/": 1, - "waitbutwhy.com": 1, - "GitHub": 1, - "https://aiimpacts.org": 0.2, - "arbital.com": 0.2, - "carado.moe": 0.3, - "nonarxiv_papers": 0.1, - "https://vkrakovna.wordpress.com": .5, - "https://jsteinhardt.wordpress.com": .5, - "audio-transcripts": 0.2, - "https://intelligence.org": .1, - "youtube": 0.07, - "reports": 0.4, - "https://aisafety.camp": 1, - "curriculum": 1, - "https://www.yudkowsky.net": 0.2, - "distill": 1, - "Cold Takes": 0.5, - "printouts": 1, - "gwern.net": 1, - "generative.ink": 1, - "greaterwrong.com": 0.2 - } - - random_number = random.random() - if random_number > desired_source_proportions[entry['source']]: - continue - - # if we specified a fraction of articles to use, only use that fraction from the remaining articles - random_number = random.random() - if random_number > self.fraction_of_articles_to_use: - continue - - # Get title, author, date, URL, tags, and text - title, author, date_published, url, tags, text = self.extract_info_from_article(entry) - - # If there's less than 2 of 'title', 'author' and 'url', ignore this text - if (((title or '').strip() == '') + ((author or '').strip() == '') + ((url or '').strip() == '')) > 1: - print(f'{entry["source"]}') - continue - - #if the text is too short, ignore this text - if len(text) < 500: - continue + # if we specified custom sources, only include entries from those sources + if (self.custom_sources is not None) and (entry['source'] not in self.custom_sources): + continue - #we're keeping the text so we inc the aticle count - self.articles_count[entry['source']] += 1 - self.total_articles_count += 1 - - # Get signature - signature = "" - if title: signature += f"Title: {title}, " - else: signature += f"Title: None, " - if author: signature += f"Author: {author}" - else: signature += f"Author: None" - # if date_published: signature += f"Date published: {date_published}, " - # if url: signature += f"URL: {url}, " - # if tags: signature += f"Tags: {tags}, " # Temporary decision to not include tags in the signature - # if signature: signature = signature[:-2] - signature = signature.replace("\n", " ") - - # Add info to metadata and embedding strings - self.metadata.append((title, author, date_published, url, tags)) - blocks = text_splitter.split(text, signature) - self.embedding_strings.extend(blocks) - self.embeddings_metadata_index.extend([self.total_articles_count-1] * len(blocks)) - - # Update counts - self.total_char_count += len(text) - self.total_word_count += len(text.split()) - self.total_sentence_count += len(split_into_sentences(text)) - self.total_block_count += len(blocks) - - except MissingDataException as e: - if str(e) not in error_count_dict: - error_count_dict[str(e)] = 0 - error_count_dict[str(e)] += 1 + # raise error if the entry does not contain the required metadata + self.contains_required_metadata(entry) + + #if the text is too short, ignore this text + if len(entry['text']) < 500: + continue + + + """Checks are done, so we construct the metadata.""" + + # Get id, source, title, text, url, date_published, authors, and summary + entry_id: str = entry['id'] + source: str = entry['source'] + title: str = entry['title'] + text: str = entry['text'] + url: str = entry['url'] + date_published: str = entry['date_published'] + authors: list = entry['authors'] + + # Get signature + if authors: + signature = f"Title: {title}, Authors: {get_authors_str(authors)}" + else: + signature = f"Title: {title}" + + # We use the text_splitter to get the chunks from the entry, + # and we add a metadata element for each new chunk we add to the dataset. + chunks = text_splitter.split(text, signature) + num_chunks = len(chunks) + + for i in range(num_chunks): + self.metadata.append({ + 'id': f"{entry_id}_{str(i+1).zfill(6)}", + 'entry_id': entry_id, + 'source': source, + 'title': title, + 'text': chunks[i], + 'url': url, + 'date_published': date_published, + 'authors': authors + }) + + # Update counts + self.entries_per_source_count[entry['source']] += 1 + self.total_counts['entries'] += 1 + self.total_counts['chars'] += len(text) + self.total_counts['words'] += len(text.split()) + self.total_counts['sentences'] += len(split_into_sentences(text)) + self.total_counts['chunks'] += len(chunks) + + def show_stats(self): + print(f'Number of entries by source: {self.entries_per_source_count}') + print(f'Total entries count: {self.total_counts["entries"]}') + print(f'Total character count: {self.total_counts["chars"]}') + print(f'Total word count: {self.total_counts["words"]}') + print(f'Total sentence count: {self.total_counts["sentences"]}') + print(f'Total chunk count: {self.total_counts["chunks"]}') def get_embeddings(self): def get_embeddings_at_index(texts: str, batch_idx: int, batch_size: int = 200): # int, np.ndarray @@ -269,15 +179,16 @@ class Dataset: rate_limit = 3500 / 60 # Maximum embeddings per second start = time.time() - self.embeddings = np.zeros((len(self.embedding_strings), LEN_EMBEDDINGS)) + embedding_strings = [chunk_data['text'] for chunk_data in self.metadata] + self.embeddings = np.zeros((len(embedding_strings), LEN_EMBEDDINGS)) with concurrent.futures.ThreadPoolExecutor() as executor: futures = [executor.submit( get_embeddings_at_index, - self.embedding_strings[batch_idx:batch_idx+batch_size], + embedding_strings[batch_idx:batch_idx+batch_size], batch_idx, - len(self.embedding_strings[batch_idx:batch_idx+batch_size]) - ) for batch_idx in range(0, len(self.embedding_strings), batch_size)] + len(embedding_strings[batch_idx:batch_idx+batch_size]) + ) for batch_idx in range(0, len(embedding_strings), batch_size)] num_completed = 0 for future in concurrent.futures.as_completed(futures): batch_idx, embeddings = future.result() @@ -289,113 +200,32 @@ class Dataset: sleep_time = max(expected_time - elapsed_time, 0) time.sleep(sleep_time) - print(f"Completed {num_completed}/{len(self.embedding_strings)} embeddings in {elapsed_time:.2f} seconds.") + print(f"Completed {num_completed}/{len(embedding_strings)} embeddings in {elapsed_time:.2f} seconds.") - def save_embeddings(self, path: str): - np.save(path, self.embeddings) - - def load_embeddings(self, path: str): - self.embeddings = np.load(path) - - def save_class(self, path: str = PATH_TO_DATASET_PKL): - # Save the class to a pickle file - print(f"Saving class to {path}...") - with open(path, 'wb') as f: - pickle.dump(self, f) - def save_data(self, path: str = PATH_TO_DATASET_DICT_PKL): # Save the data to a pickle file print(f"Saving data to {path}...") data = { "metadata": self.metadata, - "embedding_strings": self.embedding_strings, - "embeddings_metadata_index": self.embeddings_metadata_index, "embeddings": self.embeddings.astype(np.float32), - "articles_count": self.articles_count, - "total_articles_count": self.total_articles_count, - "total_char_count": self.total_char_count, - "total_word_count": self.total_word_count, - "total_sentence_count": self.total_sentence_count, - "total_block_count": self.total_block_count + "total_counts": self.total_counts, + "entries_per_source_count": self.entries_per_source_count } with open(path, 'wb') as f: pickle.dump(data, f) -def get_authors_list(authors_string: str) -> List[str]: - """ - Given a string of authors, return a list of the authors, even if the string contains a single author. - """ - authors_string = authors_string.replace(" and ", ",") - authors_string = authors_string.replace('\n', ' ') - authors = [] - if authors_string is None: - return [] - if "," in authors_string: - authors = [author.strip() for author in authors_string.split(",")] +def get_authors_str(authors_lst: List[str]) -> str: + if authors_lst == []: return 'n/a' + if len(authors_lst) == 1: return authors_lst[0] else: - authors = [authors_string.strip()] - return authors + authors_lst = authors_lst[:3] + authors_str = ", ".join(authors_lst[:-1]) + " and " + authors_lst[-1] + return authors_str def standardize_date(date_string, default_date='n/a'): try: dt = parse(date_string) return dt.strftime('%Y-%m-%d') except (ParserError, ValueError): - return default_date - - - -""" -if __name__ == "__main__": - # List of possible sources: - all_sources = ["https://aipulse.org", "ebook", "https://qualiacomputing.com", "alignment forum", "lesswrong", "manual", "arxiv", "https://deepmindsafetyresearch.medium.com", "waitbutwhy.com", "GitHub", "https://aiimpacts.org", "arbital.com", "carado.moe", "nonarxiv_papers", "https://vkrakovna.wordpress.com", "https://jsteinhardt.wordpress.com", "audio-transcripts", "https://intelligence.org", "youtube", "reports", "https://aisafety.camp", "curriculum", "https://www.yudkowsky.net", "distill", "Cold Takes", "printouts", "gwern.net", "generative.ink", "greaterwrong.com"] # These sources do not have a source field in the .jsonl file - - # List of sources we are using for the test run: - custom_sources = [ - # "https://aipulse.org", - # "ebook", - # "https://qualiacomputing.com", - # "alignment forum", - # "lesswrong", - "manual", - # "arxiv", - # "https://deepmindsafetyresearch.medium.com", - "waitbutwhy.com", - # "GitHub", - # "https://aiimpacts.org", - # "arbital.com", - # "carado.moe", - # "nonarxiv_papers", - # "https://vkrakovna.wordpress.com", - "https://jsteinhardt.wordpress.com", - # "audio-transcripts", - # "https://intelligence.org", - # "youtube", - # "reports", - "https://aisafety.camp", - "curriculum", - "https://www.yudkowsky.net", - # "distill", - # "Cold Takes", - # "printouts", - # "gwern.net", - # "generative.ink", - # "greaterwrong.com" - ] - - dataset = Dataset( - jsonl_data_path=PATH_TO_RAW_DATA.resolve(), - custom_sources=custom_sources, - rate_limit_per_minute=3500, - min_tokens_per_block=200, max_tokens_per_block=300, - # fraction_of_articles_to_use=1/2000 - ) - dataset.get_alignment_texts() - dataset.get_embeddings() - # dataset.save_embeddings("data/embeddings.npy") - - dataset.save_class(PATH_TO_DATASET.resolve()) - # # dataset = pickle.load(open("dataset.pkl", "rb")) - """ - \ No newline at end of file + return default_date \ No newline at end of file diff --git a/src/dataset/settings.py b/src/dataset/settings.py index 28cb3da..abcb211 100644 --- a/src/dataset/settings.py +++ b/src/dataset/settings.py @@ -4,9 +4,8 @@ EMBEDDING_MODEL = "text-embedding-ada-002" COMPLETIONS_MODEL = "gpt-3.5-turbo" LEN_EMBEDDINGS = 1536 -MAX_LEN_PROMPT = 4095 # This may be 8191, unsure. +MAX_LEN_PROMPT = 8190 # This may be 8191, unsure. current_file_path = Path(__file__).resolve() -PATH_TO_RAW_DATA = str(current_file_path.parent / 'data' / 'alignment_texts.jsonl') PATH_TO_DATASET_PKL = str(current_file_path.parent / 'data' / 'dataset.pkl') PATH_TO_DATASET_DICT_PKL = str(current_file_path.parent / 'data' / 'dataset_dict.pkl') \ No newline at end of file diff --git a/src/main.py b/src/main.py index 900a14f..333795a 100644 --- a/src/main.py +++ b/src/main.py @@ -27,20 +27,14 @@ if str(src_path) not in sys.path: from dataset import create_dataset #from assistant import semantic_search -from settings import EMBEDDING_MODEL, PATH_TO_DATASET_DICT_PKL +from settings import EMBEDDING_MODEL, PATH_TO_DATASET_DICT_PKL, PATH_TO_DATASET_PKL import numpy as np import matplotlib.pyplot as plt def load_rawdata_into_pkl(): - """with open(PATH_TO_DATASET, 'rb') as f: - dataset = pickle.load(f) - AS = AlignmentSearch(dataset=dataset) - prompt = "What would be an idea to solve the Alignment Problem? Name the Lesswrong post by Quintin Pope that discusses this idea." - answer = AS.search_and_answer(prompt, 3, HyDE=False) - print(answer) - """ + """ # List of possible sources: all_sources = ["https://aipulse.org", "ebook", "https://qualiacomputing.com", "alignment forum", "lesswrong", "manual", "arxiv", "https://deepmindsafetyresearch.medium.com", "waitbutwhy.com", "GitHub", "https://aiimpacts.org", "arbital.com", "carado.moe", "nonarxiv_papers", "https://vkrakovna.wordpress.com", "https://jsteinhardt.wordpress.com", "audio-transcripts", "https://intelligence.org", "youtube", "reports", "https://aisafety.camp", "curriculum", "https://www.yudkowsky.net", "distill", "Cold Takes", "printouts", "gwern.net", "generative.ink", "greaterwrong.com"] # These last do not have a source field in the .jsonl file @@ -77,87 +71,27 @@ def load_rawdata_into_pkl(): "generative.ink", "greaterwrong.com" ] - - dataset = create_dataset.Dataset( - custom_sources=custom_sources, - rate_limit_per_minute=3500, - min_tokens_per_block=200, max_tokens_per_block=300, - # fraction_of_articles_to_use=1/150, + """ + + dataset = create_dataset.ChunkedARD( + min_tokens_per_block=200, max_tokens_per_block=300 ) dataset.get_alignment_texts() - - print(len(dataset.embedding_strings)) - print(dataset.total_word_count) - print(dataset.total_block_count) - print(dataset.articles_count) - dataset.get_embeddings() dataset.save_data() - + dataset.show_stats() + +def load_pkl_and_display_stuff(): + with open(PATH_TO_DATASET_DICT_PKL, 'rb') as f: + dataset_dict = pickle.load(f) + print(dataset_dict) + @retry(wait=wait_random_exponential(min=1, max=20), stop=stop_after_attempt(4)) def get_embedding(text: str) -> np.ndarray: result = openai.Embedding.create(model=EMBEDDING_MODEL, input=text) return np.array(result["data"][0]["embedding"]) -def print_out_dataset_stuff(): - with open(PATH_TO_DATASET_PKL, 'rb') as f: - dataset = pickle.load(f) - - embeddings_len = len(dataset.embedding_strings) - i1 = random.randint(0,embeddings_len-1) - #i2 = random.randint(0,embeddings_len-1) - #print(len(dataset.embeddings)) - #print(len(dataset.embedding_strings)) - #embedding_test = get_embedding(dataset.embedding_strings[i]) - #print(np.dot(embedding_test,dataset.embeddings[i])) - - metadata_i1 = dataset.embeddings_metadata_index[i1] - print("metadata:",dataset.metadata[metadata_i1]) - print("embedding_string:",dataset.embedding_strings[i1]) - #print("embedding_vector:",dataset.embeddings[i1]) - embedding_of_string1 = get_embedding(dataset.embedding_strings[i1]) - - #metadata_i2 = dataset.embeddings_metadata_index[i2] - #print("metadata:",dataset.metadata[metadata_i2]) - #print("embedding_string:",dataset.embedding_strings[i2]) - #print("embedding_vector:",dataset.embeddings[i1]) - #embedding_of_string2 = get_embedding(dataset.embedding_strings[i2]) - #embedding_of_string2 = get_embedding("000000000000000000000000000000000000000000000000000000000000000000000000000000000") - - #print(len(embedding_of_string1)) - vector = dataset.embeddings[i1] - plot_likelihood(vector) - #plot_likelihood(get_embedding("tst")) - print(max(vector), min(vector)) - print(sum([x**2 for x in vector])) - - - - - #print(np.dot(embedding_of_string1, embedding_of_string2)) - -def plot_likelihood(embeddings, num_buckets=200): - # Calculate the histogram - histogram, bin_edges = np.histogram(embeddings.flatten(), bins=num_buckets, range=(embeddings.min(), embeddings.max())) - - # Normalize the histogram to get likelihoods - likelihoods = histogram / embeddings.flatten().size - - # Plot the likelihoods - plt.bar(bin_edges[:-1], likelihoods, width=(bin_edges[1] - bin_edges[0]), edgecolor="k", alpha=0.7) - plt.xlabel("Value") - plt.ylabel("Likelihood") - plt.title("Likelihood of Floats in the Vector Embedding") - plt.savefig("bla.png") - - - - - - if __name__ == "__main__": - # load_rawdata_into_pkl() - # print_out_dataset_stuff() - - with open(PATH_TO_DATASET_DICT_PKL, 'rb') as f: - dataset = pickle.load(f) \ No newline at end of file + load_rawdata_into_pkl() + #load_pkl_and_display_stuff() + \ No newline at end of file