#!/usr/bin/env python3

import numpy as np
import matplotlib.pyplot as plt

x = np.linspace(-6, 6, 1000)

def sigmoid(x):
    return 1 / (1 + np.exp(-x))

def silu(x):
    return x * sigmoid(x)

def gelu(x):           # standard approximation
    return 0.5 * x * (1 + np.tanh(np.sqrt(2/np.pi) * (x + 0.044715 * x**3)))

def sigmoid_prime(x):
    s = sigmoid(x)
    return s * (1 - s)

def silu_prime(x):
    s = sigmoid(x)
    return s * (1 + x * (1 - s))

def gelu_prime(x):
    # numerical derivative for the tanh-approx GELU
    h = 1e-5
    return (gelu(x + h) - gelu(x - h)) / (2 * h)

fig, axes = plt.subplots(1, 2, figsize=(12, 4.5))

# --- Left: the activations themselves ---
axes[0].plot(x, sigmoid(x), label=r'$\sigma(x)$ (sigmoid)', lw=2)
axes[0].plot(x, silu(x),   label=r'SiLU$(x)$',           lw=2)
axes[0].plot(x, gelu(x),   label=r'GELU$(x)$',           lw=2)
axes[0].axhline(0, color='k', lw=0.5)
axes[0].axvline(0, color='k', lw=0.5)
axes[0].set_title('Activations')
axes[0].set_xlabel('x')
axes[0].legend()
axes[0].grid(alpha=0.3)

# --- Right: the derivatives (gradient flow) ---
axes[1].plot(x, sigmoid_prime(x), label=r"$\sigma'(x)$",  lw=2)
axes[1].plot(x, silu_prime(x),    label=r"SiLU'$(x)$",    lw=2)
axes[1].plot(x, gelu_prime(x),    label=r"GELU'$(x)$",    lw=2)
axes[1].axhline(0, color='k', lw=0.5)
axes[1].axvline(0, color='k', lw=0.5)
axes[1].set_title('Derivatives (gradient magnitude)')
axes[1].set_xlabel('x')
axes[1].legend()
axes[1].grid(alpha=0.3)

plt.tight_layout()
plt.savefig('activation_compare.png', dpi=150)
plt.show()
