diff --git a/llm/Byte-Pair-Encoder/BPE-q3-SOLN.ipynb b/llm/Byte-Pair-Encoder/BPE-q3-SOLN.ipynb index ce3d058..e34e96e 100644 --- a/llm/Byte-Pair-Encoder/BPE-q3-SOLN.ipynb +++ b/llm/Byte-Pair-Encoder/BPE-q3-SOLN.ipynb @@ -41,54 +41,62 @@ "from collections import defaultdict, Counter\n", "\n", "def get_vocab(corpus):\n", - " \"\"\"Creates a vocabulary with words split into characters and a special end-of-word token.\"\"\"\n", - " vocab = Counter()\n", - " for word in corpus:\n", - " tokens = list(word) + ['']\n", - " vocab[tuple(tokens)] += 1\n", - " return vocab\n", + " \"\"\"Creates a vocabulary with words split into characters and a special end of word token.\"\"\"\n", + " vocab = Counter()\n", + " for word in corpus:\n", + " tokens = list(word) + [\"\"]\n", + " vocab[tuple(tokens)]+=1\n", + " return vocab\n", "\n", "def get_stats(vocab):\n", - " \"\"\"Counts frequency of adjacent symbol pairs.\"\"\"\n", - " pairs = defaultdict(int)\n", - " for word, freq in vocab.items():\n", - " for i in range(len(word) - 1):\n", - " pairs[(word[i], word[i + 1])] += freq\n", - " return pairs\n", + " \"\"\"Counts frequency of adjacent symbol pairs.\"\"\"\n", + " pairs = defaultdict(int)\n", + " for word,freq in vocab.items():\n", + " for i in range(len(word)-1):\n", + " pairs[(word[i], word[i+1])] +=freq\n", + " return pairs\n", "\n", "def merge_vocab(pair, vocab):\n", - " \"\"\"Merges the most frequent pair into a single symbol.\"\"\"\n", - " new_vocab = {}\n", - " bigram = ' '.join(pair)\n", - " replacement = ''.join(pair)\n", - " for word, freq in vocab.items():\n", - " word_str = ' '.join(word)\n", - " # Replace bigram with merged symbol\n", - " new_word_str = word_str.replace(bigram, replacement)\n", - " new_vocab[tuple(new_word_str.split())] = freq\n", - " return new_vocab\n", + " \"\"\"Merges the most frequent pair into a single symbol.\"\"\"\n", + " new_vocab = {}\n", + " bigram = ' '.join(pair)\n", + " replacement = ''.join(pair)\n", + " print(f\"Bigram: {bigram}\")\n", + " for word,freq in vocab.items():\n", + " word_str = ' '.join(word)\n", + " #Replace bigram with merged symbol\n", + " new_word_str=word_str.replace(bigram, replacement)\n", + " #Preserve the trailing empty string token if the original word ended with it\n", + " tokens = new_word_str.split()\n", + " if word[-1] == \"\":\n", + " tokens.append(\"\")\n", + " new_vocab[tuple(tokens)]=freq\n", + " return new_vocab\n", "\n", - "def byte_pair_encoding(corpus, num_merges=10):\n", - " \"\"\"Performs BPE on a corpus.\"\"\"\n", - " vocab = get_vocab(corpus)\n", - " merges = []\n", - " for _ in range(num_merges):\n", - " pairs = get_stats(vocab)\n", - " if not pairs:\n", - " break\n", - " best = max(pairs, key=pairs.get)\n", - " vocab = merge_vocab(best, vocab)\n", - " merges.append(best)\n", - " print(f\"Merge {_ + 1}: {best}\")\n", - " return vocab, merges\n", "\n", - "# Example usage\n", - "corpus = [\"low\", \"lowest\", \"newer\", \"wider\"]\n", - "final_vocab, merge_operations = byte_pair_encoding(corpus, num_merges=10)\n", + "def byte_pair_encoding(corpus, num_merges=10):\n", + " \"\"\"Performs BPE on a corpus.\"\"\"\n", + " vocab = get_vocab(corpus)\n", + " merges = []\n", + " for _ in range(num_merges):\n", + " pairs = get_stats(vocab)\n", + " if not pairs:\n", + " break\n", + " best = max(pairs, key=pairs.get)\n", + " vocab = merge_vocab(best, vocab)\n", + " merges.append(best)\n", + " print(f\"Merge {_ + 1}: {best}\")\n", + " return vocab, merges\n", "\n", + "corpus = [\"low\", \"lowest\", \"newer\", \"wider\", \"owl\", \"LOW\"]\n", + "final_vocab , merge_operations = byte_pair_encoding(corpus, num_merges=10)\n", + "# vocab = get_vocab(corpus)\n", + "# pairs = get_stats(vocab)\n", + "# print(pairs)\n", + "# new_vocab = merge_vocab(pairs, vocab)\n", "print(\"\\nFinal Vocabulary:\")\n", "for word in final_vocab:\n", - " print(' '.join(word), \":\", final_vocab[word])\n" + " print(\"\".join(word), \":\", final_vocab[word])" ] }, { @@ -99,41 +107,43 @@ "outputs": [], "source": [ "def test_get_vocab():\n", - " corpus = [\"test\"]\n", - " vocab = get_vocab(corpus)\n", - " assert vocab == {('t', 'e', 's', 't', ''): 1}\n", - " print(\"✓ test_get_vocab passed\")\n", + " corpus = ['test']\n", + " vocab = get_vocab(corpus)\n", + " assert vocab == {(\"t\", \"e\", \"s\", \"t\", \"\"):1}\n", + " print(\"✓\")\n", "\n", - "def test_get_stats():\n", - " vocab = {('t', 'e', 's', 't', ''): 1}\n", - " stats = get_stats(vocab)\n", - " expected = {\n", + "def test_get_status():\n", + " vocab = {('t', 'e', 's', 't', '' ):1}\n", + " stats = get_stats(vocab)\n", + " expected = {\n", " ('t', 'e'): 1,\n", " ('e', 's'): 1,\n", " ('s', 't'): 1,\n", - " ('t', ''): 1\n", - " }\n", - " assert stats == expected\n", - " print(\"✓ test_get_stats passed\")\n", + " ('t', ''): 1\n", + " }\n", + " assert stats == expected\n", + " print(\"✓✓\")\n", "\n", "def test_merge_vocab():\n", - " vocab = {('t', 'e', 's', 't', ''): 1}\n", - " merged = merge_vocab(('e', 's'), vocab)\n", - " expected = {('t', 'es', 't', ''): 1}\n", - " assert merged == expected\n", - " print(\"✓ test_merge_vocab passed\")\n", + " vocab = {('t', 'e', 's', 't', '' ):1}\n", + " merged = merge_vocab((\"e\", \"s\"),vocab)\n", + " print(merged)\n", + " expected = {(\"t\", \"es\", 't', ''):1}\n", + " assert merged == expected\n", + " print(\"✓✓✓\")\n", "\n", "def test_bpe_sequence():\n", - " corpus = [\"low\", \"lower\", \"newest\", \"widest\"]\n", - " final_vocab, merges = byte_pair_encoding(corpus, num_merges=5)\n", - " assert isinstance(final_vocab, dict)\n", - " assert all(isinstance(pair, tuple) for pair in merges)\n", - " assert len(merges) == 5\n", - " print(\"✓ test_bpe_sequence passed\")\n", + " corpus = ['low', 'lower', 'lower', 'newest', 'wildest']\n", + " final_vocab, merges = byte_pair_encoding(corpus, num_merges=5)\n", + " assert isinstance(final_vocab, dict)\n", + " assert all(isinstance(pair, tuple) for pair in merges)\n", + " assert len(merges) == 5\n", + " print(\"✓✓✓✓\")\n", "\n", - "# Run all tests\n", "test_get_vocab()\n", - "test_get_\n" + "test_get_status()\n", + "test_merge_vocab()\n", + "test_bpe_sequence()" ] } ], diff --git a/website/src/data/questions.ts b/website/src/data/questions.ts index 5e5437a..aa7551e 100644 --- a/website/src/data/questions.ts +++ b/website/src/data/questions.ts @@ -441,7 +441,7 @@ const v2Questions: Omit