{
  "nbformat": 4,
  "nbformat_minor": 0,
  "metadata": {
    "colab": {
      "provenance": []
    },
    "kernelspec": {
      "name": "python3",
      "display_name": "Python 3"
    },
    "language_info": {
      "name": "python"
    }
  },
  "cells": [
    {
      "cell_type": "code",
      "execution_count": 2,
      "metadata": {
        "colab": {
          "base_uri": "https://localhost:8080/"
        },
        "id": "JCROI8pBalHB",
        "outputId": "09ec4f45-5bb1-4336-dd08-0398017ca641"
      },
      "outputs": [
        {
          "output_type": "stream",
          "name": "stderr",
          "text": [
            "[nltk_data] Downloading package reuters to /root/nltk_data...\n",
            "[nltk_data]   Package reuters is already up-to-date!\n",
            "[nltk_data] Downloading package punkt to /root/nltk_data...\n",
            "[nltk_data]   Package punkt is already up-to-date!\n",
            "[nltk_data] Downloading package punkt_tab to /root/nltk_data...\n",
            "[nltk_data]   Unzipping tokenizers/punkt_tab.zip.\n"
          ]
        },
        {
          "output_type": "stream",
          "name": "stdout",
          "text": [
            "Next Word: of\n"
          ]
        }
      ],
      "source": [
        "import nltk\n",
        "from nltk import trigrams\n",
        "from nltk.corpus import reuters\n",
        "from collections import defaultdict\n",
        "\n",
        "nltk.download('reuters')\n",
        "nltk.download('punkt')\n",
        "nltk.download('punkt_tab')\n",
        "\n",
        "words = nltk.word_tokenize(' '.join(reuters.words()))\n",
        "tri_grams = list(trigrams(words))\n",
        "\n",
        "model = defaultdict(lambda: defaultdict(lambda: 0))\n",
        "\n",
        "for w1, w2, w3 in tri_grams:\n",
        "    model[(w1, w2)][w3] += 1\n",
        "\n",
        "for w1_w2 in model:\n",
        "    total_count = float(sum(model[w1_w2].values()))\n",
        "\n",
        "    for w3 in model[w1_w2]:\n",
        "        model[w1_w2][w3] /= total_count\n",
        "\n",
        "\n",
        "def predict_next_word(w1, w2):\n",
        "    next_word_probs = model[(w1, w2)]\n",
        "\n",
        "    if next_word_probs:\n",
        "        return max(next_word_probs, key=next_word_probs.get)\n",
        "    else:\n",
        "        return \"No prediction available\"\n",
        "\n",
        "\n",
        "print(\"Next Word:\", predict_next_word('the', 'stock'))"
      ]
    },
    {
      "cell_type": "code",
      "source": [],
      "metadata": {
        "id": "ma67yM-ual5Q"
      },
      "execution_count": null,
      "outputs": []
    }
  ]
}