"""Reproduce the tiny, illustrative calculations in the accompanying NLP blog post.

Run from the repository root: python assets/ai/nlp_examples.py
Only Python's standard library is used. No files or network are accessed.
These examples are arithmetic demonstrations, not trained NLP models.
"""

import math


def dot(left, right):
    """Multiply corresponding coordinates, then sum."""
    return sum(a * b for a, b in zip(left, right))


def sigmoid(score):
    return 1.0 / (1.0 + math.exp(-score))


def softmax(scores):
    """Subtract the maximum for numerical stability; -inf masks are supported."""
    maximum = max(scores)
    exponentials = [math.exp(score - maximum) for score in scores]
    total = sum(exponentials)
    return [value / total for value in exponentials]


def weighted_sum(weights, values):
    return [
        sum(weight * row[column] for weight, row in zip(weights, values))
        for column in range(len(values[0]))
    ]


def rounded(values):
    return [round(value, 6) for value in values]


def show(title, **values):
    print(f"\n{title}")
    for name, value in values.items():
        print(f"  {name}: {value}")


def main():
    # English aliases preserve the seven-column vocabulary order in chapter 2.
    vocabulary = ["movie", "plot", "good", "exciting", "bad", "boring", "not"]
    review = ["movie", "not", "good"]
    bow = [review.count(word) for word in vocabulary]
    assert bow == [1, 0, 1, 0, 0, 0, 1]
    show("Ch. 3: bag of words", vocabulary=vocabulary, vector=bow)

    movie_idf = math.log(3 / 3)
    good_idf = math.log(3 / 1)
    show("Ch. 3: TF-IDF", movie_idf=movie_idf, good_idf=round(good_idf, 6))

    # One example, one weight, no bias: x=1, y=1, learning rate=0.2.
    weight, feature, target, learning_rate = 0.0, 1.0, 1.0, 0.2
    before = sigmoid(weight * feature)
    gradient = (before - target) * feature
    weight -= learning_rate * gradient
    after = sigmoid(weight * feature)
    assert after > before
    show(
        "Ch. 4: one classification update",
        probability_before=before,
        loss_before=round(-math.log(before), 6),
        gradient=gradient,
        weight_after=weight,
        probability_after=round(after, 6),
        loss_after=round(-math.log(after), 6),
    )

    show("Ch. 5: add-one smoothing", probabilities=rounded([3 / 6, 2 / 6, 1 / 6]))

    sentence_mean = [0.5, 0.5]
    score = dot([2, 0], sentence_mean)
    show("Ch. 6: mean embedding classification", mean=sentence_mean, score=score,
         probability=round(sigmoid(score), 6))

    # CBOW: [like, movie] predicts watch. Candidate order: watch, eat, buy.
    inputs = [[1.0, 0.0], [0.0, 1.0]]
    outputs = [[2.0, 2.0], [0.0, 0.0], [1.0, 1.0]]
    mean = weighted_sum([0.5, 0.5], inputs)
    scores = [dot(mean, output) for output in outputs]
    probabilities = softmax(scores)
    score_gradients = probabilities.copy()
    score_gradients[0] -= 1.0
    mean_gradient = weighted_sum(score_gradients, outputs)
    input_gradient = [value / 2 for value in mean_gradient]
    updated_inputs = [
        [value - 0.1 * grad for value, grad in zip(row, input_gradient)]
        for row in inputs
    ]
    # Freeze output vectors here to isolate the input-side update, as in the notes.
    updated_mean = weighted_sum([0.5, 0.5], updated_inputs)
    updated_probabilities = softmax([dot(updated_mean, row) for row in outputs])
    assert updated_probabilities[0] > probabilities[0]
    show(
        "Ch. 7: CBOW and an input-embedding update",
        scores=scores,
        probabilities=rounded(probabilities),
        loss=round(-math.log(probabilities[0]), 6),
        mean_gradient=rounded(mean_gradient),
        input_gradient=rounded(input_gradient),
        updated_inputs=[rounded(row) for row in updated_inputs],
        probability_after_input_update=round(updated_probabilities[0], 6),
    )

    negative_sampling_loss = -math.log(sigmoid(2)) - math.log(sigmoid(1))
    show("Ch. 7: one positive and one noise pair", loss=round(negative_sampling_loss, 6))

    # Deliberately simplified linear recurrence: no tanh, not a complete RNN.
    def trace(sequence):
        state = 0.0
        history = []
        for value in sequence:
            state = 0.5 * state + value
            history.append(state)
        return history

    forward = trace([1, 0, 2])
    reverse = trace([2, 0, 1])
    assert forward[-1] == 2.25 and reverse[-1] == 1.5
    show("Ch. 8: order-sensitive recurrence", cat_chases_dog=forward,
         dog_chases_cat=reverse, retention_20=0.5 ** 20)

    memory = 0.9 * 0.8 + 0.2 * 0.5
    hidden = 0.7 * math.tanh(memory)
    show("Ch. 9: one LSTM memory coordinate", memory=round(memory, 6),
         hidden_output=round(hidden, 6))

    show("Ch. 10: one convolution window", original=dot([2, 1, 1], [-1, 0.5, 1]),
         reversed=dot([2, 1, 1], [1, 0.5, -1]))

    # Both top-k=2 and top-p=0.8 select the first two entries for this distribution.
    distribution = [0.6, 0.3, 0.1]
    top_two = [p / sum(distribution[:2]) for p in distribution[:2]]
    truth_probabilities = [0.5, 0.25, 0.125]
    average_loss = -sum(math.log(p) for p in truth_probabilities) / 3
    perplexity = math.exp(average_loss)
    assert math.isclose(perplexity, 4)
    show("Ch. 11: decoding and perplexity", top_two=rounded(top_two),
         average_loss=round(average_loss, 6), perplexity=round(perplexity, 6))

    cross_weights = softmax([0, 1, 2])
    cross_context = weighted_sum(cross_weights, [[1, 0], [0, 1], [1, 1]])
    show("Ch. 12: translation attention", weights=rounded(cross_weights),
         context=rounded(cross_context))

    # Q=K=X, while W_V doubles the second feature coordinate.
    query = [1, 0]
    keys = [[1, 0], [0, 1], [1, 1]]
    values = [[1, 0], [0, 2], [1, 2]]
    attention_scores = [dot(query, key) / math.sqrt(2) for key in keys]
    weights = softmax(attention_scores)
    output = weighted_sum(weights, values)
    masked_weights = softmax([attention_scores[0], -math.inf, -math.inf])
    masked_output = weighted_sum(masked_weights, values)
    assert math.isclose(sum(weights), 1) and masked_weights == [1.0, 0.0, 0.0]
    assert masked_output == [1.0, 0.0]
    show("Ch. 13: self-attention, first query", scaled_scores=rounded(attention_scores),
         weights=rounded(weights), output=rounded(output),
         causal_weights=masked_weights, causal_output=masked_output)

    embedding_bytes = 10_000 * 128 * 4
    attention_bytes_one_head = 4096 * 4096 * 2
    kv_bytes = 2 * 1 * 32 * 8192 * 8 * 128 * 2
    assert kv_bytes == 1024 ** 3
    show("Ch. 14: illustrative memory sizes", embedding_MB=embedding_bytes / 1_000_000,
         attention_one_head_MiB=attention_bytes_one_head / 1024 ** 2,
         attention_32_heads_GiB=attention_bytes_one_head * 32 / 1024 ** 3,
         kv_cache_GiB=kv_bytes / 1024 ** 3)

    precision, recall = 8 / 10, 8 / 12
    f1 = 2 * precision * recall / (precision + recall)
    show("Ch. 16: entity metrics", precision=precision, recall=round(recall, 6),
         f1=round(f1, 6), crf_complete_entity_score=2 + 1.5 + 1,
         crf_partial_entity_score=2 + 1.8)

    full_parameters, lora_parameters = 4096 ** 2, 8 * (4096 + 4096)
    show("Ch. 17: one matrix with LoRA", full_parameters=full_parameters,
         lora_trainable_parameters=lora_parameters,
         ratio=full_parameters / lora_parameters)

    show("Ch. 18: padding and mean pooling", correct_mean=[0.5, 0.5],
         incorrect_mean=rounded([1 / 3, 1 / 3]))
    show("Appendix A: one-word Bayes example", positive_probability=0.2 / (0.2 + 0.01))

    print("\nAll numerical checks passed. These are illustrative calculations, not model benchmarks.")


if __name__ == "__main__":
    main()
