"""Count-based next-word prediction. Python 3.10+, standard library only.

The built-in corpus is a public-domain excerpt of Shakespeare's Hamlet,
Act III, scene i. Lowercase words are tokens; punctuation is discarded.
We treat the excerpt as one sequence and add an end marker, not line breaks.
This is a teaching example, not an LLM tokenizer or a useful Shakespeare model.
"""
import argparse
from collections import Counter, defaultdict
from pathlib import Path
import random
import re

CORPUS = """To be, or not to be: that is the question:
Whether 'tis nobler in the mind to suffer
The slings and arrows of outrageous fortune,
Or to take arms against a sea of troubles,
And by opposing end them? To die: to sleep;
No more; and by a sleep to say we end
The heart-ache and the thousand natural shocks
That flesh is heir to, 'tis a consummation
Devoutly to be wish'd. To die, to sleep;
To sleep: perchance to dream: ay, there's the rub;
For in that sleep of death what dreams may come
When we have shuffled off this mortal coil,
Must give us pause."""
END = "<end>"
START = "<start>"


def tokenize(text):
    return re.findall(r"[a-z]+(?:'[a-z]+)?", text.lower())


class NGram:
    def __init__(self, text, order):
        if order not in (1, 2, 3):
            raise ValueError("Choose order 1, 2, or 3")
        tokens = tokenize(text)
        if not tokens:
            raise ValueError("Corpus must contain English words")
        self.order = order
        self.vocabulary = set(tokens) | {END}
        self.counts = defaultdict(Counter)
        padded = [START] * (order - 1) + tokens + [END]
        for i in range(order - 1, len(padded)):
            context = tuple(padded[i - order + 1:i])
            self.counts[context][padded[i]] += 1

    def distribution(self, history):
        # Pad short histories just as at the beginning of the training text.
        width = self.order - 1
        context = tuple(([START] * width + list(history))[-width:]) if width else ()
        counts = self.counts.get(context, {})
        total = sum(counts.values())
        # Missing entries have zero mass. An unseen context has NO estimate.
        return {word: count / total for word, count in sorted(counts.items())} if total else {}

    def generate(self, prompt, steps=15, seed=7):
        rng = random.Random(seed)
        history = tokenize(prompt)
        for _ in range(steps):
            distribution = self.distribution(history)
            if not distribution:
                return " ".join(history), "unseen context: no estimate (no fallback)"
            word = rng.choices(list(distribution), weights=list(distribution.values()))[0]
            if word == END:
                return " ".join(history), "end marker selected"
            history.append(word)
        return " ".join(history), "step limit reached"


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--corpus", type=Path, help="Optional UTF-8 plain-text corpus")
    parser.add_argument("--prompt", default="to be")
    parser.add_argument("--steps", type=int, default=15, choices=range(1, 101), metavar="1..100")
    parser.add_argument("--seed", type=int, default=7)
    args = parser.parse_args()
    text = args.corpus.read_text(encoding="utf-8") if args.corpus else CORPUS
    for order in (1, 2, 3):
        model = NGram(text, order)
        print(f"\n{order}-gram: use {order - 1} previous words")
        distribution = model.distribution(tokenize(args.prompt))
        print(f"Next-word probabilities after {args.prompt!r}:")
        for word, probability in sorted(distribution.items(), key=lambda pair: (-pair[1], pair[0])):
            print(f"  {word:14} {probability:7.2%}")
        print("Other vocabulary entries: 0%" if distribution else "Unseen context: no estimate")
        print("Generated:", *model.generate(args.prompt, args.steps, args.seed), sep="\n  ")
        print("Observed stored continuations:", sum(map(len, model.counts.values())))
        print("Vocabulary size (includes end marker):", len(model.vocabulary))
    print("\nHypothetical dense table, 50,000-word vocabulary (not allocated):")
    for order in (1, 2, 3):
        cells = 50_000 ** order
        print(f"  {order}-gram: {cells:,} cells; {cells * 4:,} bytes at float32")


if __name__ == "__main__":
    main()
