"""Proof: the stability patterns. exp overflows at 88.73 in float32
and 11.09 in float16, so a naive softmax dies on logits your model
produces every day; subtracting the max moves nothing and fixes it
(softmax is shift-invariant)."""
import torch

print(f"float32 exp cliff: exp(88.7)={torch.tensor(88.7).exp().item():.3e}, "
      f"exp(88.8)={torch.tensor(88.8).exp().item()}")
h = torch.tensor([11.0, 11.1], dtype=torch.float16)
print(f"float16 exp cliff: exp(11.0)={h.exp()[0].item()}, exp(11.1)={h.exp()[1].item()}")

logits = torch.tensor([120.0, 119.0, 115.0])
naive = logits.exp() / logits.exp().sum()
stable = (logits - logits.max()).exp() / (logits - logits.max()).exp().sum()
print(f"naive  softmax([120,119,115]) = {naive.tolist()}")
print(f"stable softmax([120,119,115]) = {[round(v,4) for v in stable.tolist()]}")
print(f"torch.softmax agrees with stable: "
      f"{torch.allclose(torch.softmax(logits, 0), stable)}")

# logsumexp, same illness, same cure
print(f"naive  log(sum(exp)) = {logits.exp().sum().log().item()}")
print(f"stable logsumexp     = {torch.logsumexp(logits, 0).item():.4f}")
