# Exported from sorting-as-gradient-flow_paper-figures.ipynb
# Generated by direct notebook code-cell extraction because local nbconvert is blocked by a missing extension.

# %% [markdown]
# #### Sorting as Gradient Flow on the Permutohedron
# - Jonathan Landers, 2025
#
# This notebook contains code to generate the figures in the paper.

# %% [markdown]
# #### setup

# %%
import warnings
import matplotlib
import matplotlib.pyplot as plt

warnings.filterwarnings('ignore')

# Method 1: More specific warning filter
warnings.filterwarnings("ignore",
                       message=".*The PostScript backend does not support transparency.*",
                       category=UserWarning)

# Method 2: Alternative way to target matplotlib warnings
matplotlib.set_loglevel("critical")  # Only show critical errors, ignore warnings
matplotlib.rcParams.update({
    "pdf.fonttype": 42,
    "ps.fonttype": 42,
})

# Method 3: Using a regular expression pattern
import re
warnings.filterwarnings("ignore",
                       message=str(re.compile("The PostScript backend does not support transparency.*")))

# %%
# import fitz        # PyMuPDF
# from PIL import Image
# import io

# def pdf_to_eps(input_pdf, output_eps, page_num=0, dpi=300):
#     # Open the PDF with PyMuPDF (no EOF marker issues)
#     doc = fitz.open(input_pdf)
#     page = doc.load_page(page_num)

#     # Render at the desired DPI
#     zoom = dpi / 72  # PDF default 72 dpi
#     mat = fitz.Matrix(zoom, zoom)
#     pix = page.get_pixmap(matrix=mat, alpha=False)

#     # Load into Pillow via in-memory PPM
#     img_data = pix.tobytes("ppm")
#     img = Image.open(io.BytesIO(img_data))

#     # Save out as EPS at the correct DPI
#     img.save(output_eps, format="EPS", dpi=(dpi, dpi))
#     print(f"âœ“ Saved {output_eps!r} at {dpi} DPI")


# %% [markdown]
# #### decision tree perspective

# %%
import matplotlib.pyplot as plt
from matplotlib import rcParams

# WITH EXAMPLE TRAVERSAL

# ===== GLOBAL STYLE (matches your other figure) =====
fsize = 16
rcParams.update({
    'font.family': 'Times New Roman',
    'font.size': fsize,
    'axes.titlesize': fsize,
    'axes.labelsize': fsize,
    'xtick.labelsize': fsize,
    'ytick.labelsize': fsize,
    'legend.fontsize': fsize,
    'lines.linewidth': 1
})

# ===== COLOR SCHEME =====
# MAIN_BLUE       = '#4A6E8A'
# SECONDARY_BLUE  = '#A4B8C9'
# HIGHLIGHT_COLOR = '#2C3E50'
# GRAY_DARK       = '#5A5A5A'
# BLACK           = '#000000'
# WHITE           = '#FFFFFF'

MAIN_BLUE       = '#4A6E8A'
SECONDARY_BLUE  = '#A4B8C9'
HIGHLIGHT_COLOR = '#2C3E50'
GRAY_DARK       = '#5A5A5A'
BLACK           = '#000000'
WHITE           = '#FFFFFF'

# ===== FIXED LEVELS (perfectly even vertically) =====
y_root  = 0.95   # depth 0
y_lvl1  = 0.75   # depth 1
y_lvl2  = 0.55   # depth 2
y_leaf  = 0.35   # leaves

#fig, ax = plt.subplots(figsize=(7.2, 4.2))
fig, ax = plt.subplots(figsize=(10, 6))
ax.set_xlim(0.0, 1.0)
ax.set_ylim(0.25, 1.02)
ax.axis("off")

# ---------- Node / leaf coordinates ----------
nodes = {
    "root": (0.50, y_root),   # â„“1 < â„“2
    "n1L":  (0.25, y_lvl1),   # â„“2 < â„“3
    "n1R":  (0.75, y_lvl1),   # â„“1 < â„“3
    "n2L":  (0.25, y_lvl2),   # â„“1 < â„“3
    "n2R":  (0.75, y_lvl2),   # â„“2 < â„“3
}

leaves = {
    "L123": (0.15, y_leaf),
    "L132": (0.25, y_leaf),
    "L312": (0.35, y_leaf),
    "L213": (0.65, y_leaf),
    "L231": (0.75, y_leaf),
    "L321": (0.85, y_leaf),
}

leaf_labels = {
    "L123": r"$\sigma = 1\,2\,3$",
    "L132": r"$\sigma = 1\,3\,2$",
    "L312": r"$\sigma = 3\,1\,2$",
    "L213": r"$\sigma = 2\,1\,3$",
    "L231": r"$\sigma = 2\,3\,1$",
    "L321": r"$\sigma = 3\,2\,1$" + "\n" + r"sort output values $= [2, 4, 5]$",
}

# ---------- Edge list (no manual offsets; weâ€™ll compute them) ----------
edges = [
    ("root", "n1L",  "YES"),
    ("root", "n1R",  "NO"),
    ("n1L",  "L123", "YES"),
    ("n1L",  "n2L",  "NO"),
    ("n2L",  "L132", "YES"),
    ("n2L",  "L312", "NO"),
    ("n1R",  "L213", "YES"),
    ("n1R",  "n2R",  "NO"),
    ("n2R",  "L231", "YES"),
    ("n2R",  "L321", "NO"),
]

def get_xy(name):
    return nodes[name] if name in nodes else leaves[name]

# ---------- Draw edges + YES/NO labels slightly off the lines ----------
for parent, child, label in edges:
    x0, y0 = get_xy(parent)
    x1, y1 = get_xy(child)

    # edge
    if parent == "root" and child == "n1R":
        ax.plot([x0, x1], [y0, y1],
            color=BLACK, linewidth=5)
    elif parent == "n1R" and child == "n2R":
        ax.plot([x0, x1], [y0, y1],
            color=BLACK, linewidth=5)
    elif parent == "n2R" and child == "L321":
        ax.plot([x0, x1], [y0, y1],
            color=BLACK, linewidth=5)
    else:
        ax.plot([x0, x1], [y0, y1],
                color=GRAY_DARK, linewidth=1.5)

    # midpoint
    xm = 0.5 * (x0 + x1)
    ym = 0.5 * (y0 + y1)

    # small perpendicular offset so text is not on top of the line
    vx, vy = (x1 - x0), (y1 - y0)
    # normal vector (vy, -vx)
    norm_len = (vx**2 + vy**2) ** 0.5
    if norm_len == 0:
        nx, ny = 0.0, 0.0
    else:
        nx, ny = vy / norm_len, -vx / norm_len

    # move YES and NO to opposite sides of the edge
    side = 1.0 if label == "YES" else -1.0
    offset = 0.03  # magnitude of offset
    lx = xm + side * offset * nx
    ly = ym + side * offset * ny

    ax.text(lx, ly, label,
            ha="center", va="center",
            color=BLACK, fontsize=fsize, fontweight="bold")

# ---------- Draw nodes ----------
for x, y in nodes.values():
    ax.scatter(x, y, s=130,
               facecolor=WHITE,
               edgecolor=HIGHLIGHT_COLOR,
               linewidth=1.5,
               zorder=3)

for x, y in leaves.values():
    ax.scatter(x, y, s=130,
               facecolor=SECONDARY_BLUE,
               edgecolor=HIGHLIGHT_COLOR,
               linewidth=1.0,
               zorder=3)

# ---------- Labels on nodes ----------
# root stays exactly where it was
x_root, y_root_plot = nodes["root"]
ax.text(x_root, y_root_plot + .01, r"$\ell_1 < \ell_2$",
        ha="center", va="bottom",
        color=BLACK, fontsize=fsize)

# level 1: shift both labels toward the center
x, y = nodes["n1L"]
ax.text(x + 0.03, y, r"$\ell_2 < \ell_3$",
        ha="left", va="center",
        color=BLACK, fontsize=fsize)

x, y = nodes["n1R"]
ax.text(x + .03, y, r"$\ell_1 < \ell_3$",
        ha="left", va="center",
        color=BLACK, fontsize=fsize)

# level 2: â„“1 < â„“3 toward center; â„“2 < â„“3 pushed outward/right
x, y = nodes["n2L"]
ax.text(x + 0.03, y, r"$\ell_1 < \ell_3$",
        ha="left", va="center",
        color=BLACK, fontsize=fsize)

x, y = nodes["n2R"]
ax.text(x + 0.03, y, r"$\ell_2 < \ell_3$",
        ha="left", va="center",
        color=BLACK, fontsize=fsize)

# ---------- Leaf labels ----------
for name, label in leaf_labels.items():
    x, y = leaves[name]
    leaf_fontsize = 15 if name == "L321" else 14
    ax.text(x, y - 0.03, label,
            ha="center", va="top",
            color=BLACK, fontsize=leaf_fontsize)

# ---------- Annotations: root arrow & height ----------
# ax.annotate(
#     "root",
#     xy=nodes["root"],
#     xytext=(0.68, 0.97),
#     fontsize=12,
#     color=MAIN_BLUE,
#     arrowprops=dict(arrowstyle="->",
#                     color=MAIN_BLUE,
#                     linewidth=1.4),
# )

ax.plot([0.95, 0.95], [y_leaf, y_root],
        color=HIGHLIGHT_COLOR, linewidth=2)
ax.text(0.97, 0.5 * (y_root + y_leaf),
        r"height = 3",
        rotation=90,
        va="center", ha="center",
        color=BLACK, fontsize=17)

ax.set_title(r"sort input values $L = [\ell_1 = 5,\ell_2 = 4,\ell_3 = 2]$",
             color=BLACK, fontsize=17)

plt.tight_layout()
plt.savefig("decision_tree_l_sort_tweaked.png", dpi=300, bbox_inches="tight")
plt.savefig("decision_tree_l_sort_tweaked.pdf", dpi=300, bbox_inches="tight")
plt.savefig("decision_tree_l_sort_tweaked.eps", dpi=300, bbox_inches="tight")
plt.show()

from PyPDF2 import PdfReader, PdfWriter

input_pdf  = "decision_tree_l_sort_tweaked.pdf"
output_pdf = "decision_tree_l_sort_tweaked.pdf"

# Target size constraints (AMM or other)
target_width_inch  = 5.0   # max width in inches
target_height_inch = 8.0   # max height in inches

# Convert inches to PDF points (1 in = 72 pt)
target_w_pt = target_width_inch  * 72
target_h_pt = target_height_inch * 72

reader = PdfReader(input_pdf)
writer = PdfWriter()
page   = reader.pages[0]

# Original page size in points
orig_w_pt = float(page.mediabox.width)
orig_h_pt = float(page.mediabox.height)

# Compute uniform scale to fit within box
scale_w = target_w_pt / orig_w_pt
scale_h = target_h_pt / orig_h_pt
scale   = min(scale_w, scale_h)

# Scale page vectorially
page.scale_by(scale)

# Write out the new PDF
writer.add_page(page)
with open(output_pdf, "wb") as f:
    writer.write(f)

print(f"Saved '{output_pdf}' scaled to fit within {target_width_inch}\"Ã—{target_height_inch}\" inches (vector-only).")

# %% [markdown]
# #### geometric perspective and the permutohedron

# %%
import numpy as np
import matplotlib.pyplot as plt
from matplotlib import rcParams

fsize = 18

# ===== GLOBAL STYLE (match decision_tree_bubble_style) =====
rcParams.update({
    'font.family': 'Times New Roman',
    'font.size': fsize,
    'axes.titlesize': fsize,
    'axes.labelsize': fsize,
    'xtick.labelsize': fsize,
    'ytick.labelsize': fsize,
    'legend.fontsize': fsize,
    'lines.linewidth': 1
})

# ===== COLOR SCHEME (same as tree figure) =====
BLACK           = '#000000'
MAIN_BLUE       = BLACK
SECONDARY_BLUE  = '#A4B8C9'
HIGHLIGHT_COLOR = BLACK
GRAY_DARK       = '#5A5A5A'
WHITE           = '#FFFFFF'

# ===== SINGLE FIGURE: Diagram 4 â€“ Walk on the S3 permutohedron =====
fig, ax4 = plt.subplots(figsize=(6, 6))

# Regular hexagon for the six permutations of S3
angles = np.linspace(0, 2 * np.pi, 6, endpoint=False)
hx = np.cos(angles)
hy = np.sin(angles)

# Order the permutations around the hexagon
perm_labels = ['123', '132', '312', '321', '231', '213']
perm_to_idx = {p: i for i, p in enumerate(perm_labels)}

# Pretty TeX labels
perm_tex = {
    '123': "$v_s$\n$(1,2,3)$",
    '132': r"$(1,3,2)$",
    '312': r"$(3,1,2)$",
    '321': "$x_0$\n$(3,2,1)$",
    '231': r"$(2,3,1)$",
    '213': r"$(2,1,3)$",
}

# ----- Draw hexagon edges (all permutations and adjacency) -----
for i in range(6):
    j = (i + 1) % 6
    ax4.plot(
        [hx[i], hx[j]],
        [hy[i], hy[j]],
        color=GRAY_DARK,
        linewidth=1.5,
        zorder=1
    )

# ----- Draw vertices -----
ax4.scatter(
    hx, hy,
    s=80,
    facecolor=SECONDARY_BLUE,
    edgecolor=HIGHLIGHT_COLOR,
    linewidth=1.0,
    zorder=3
)

# ----- Permutation labels (with manual tweaks for 321 and 123) -----
for i, p in enumerate(perm_labels):
    x = 1.15 * hx[i]
    y = 1.15 * hy[i]
    ha = "center"

    if p == '321':
        x = -1.32

    if p == '123':
        x = 1.32

    if p == '231':
        y -= 0.05

    if p == '213':
        y -= 0.05

    ax4.text(
        x,
        y,
        perm_tex[p],
        ha=ha,
        va="center",
        color=BLACK,
        fontsize=fsize
    )

# ----- Highlighted path: 321 â†’ 231 â†’ 213 â†’ 123 -----
path_perm = ['321', '231', '213', '123']
constraints_perm = [
    r"$x_1>x_2$",
    r"$x_2>x_3$",
    r"$x_1>x_2$",
]

for i in range(len(path_perm) - 1):
    a = perm_to_idx[path_perm[i]]
    b = perm_to_idx[path_perm[i + 1]]

    # thick highlighted edge
    ax4.plot(
        [hx[a], hx[b]],
        [hy[a], hy[b]],
        color=HIGHLIGHT_COLOR,
        linewidth=3.5,
        zorder=2
    )

    # constraint label near midpoint, nudged outward a bit
    xm = 0.5 * (hx[a] + hx[b])
    ym = 0.5 * (hy[a] + hy[b])

    # small radial push away from origin so labels don't sit on edges
    r = (xm**2 + ym**2) ** 0.9
    if r != 0:
        xm_off = xm * 1.1 / r
        ym_off = ym * 1.1 / r
    else:
        xm_off, ym_off = xm, ym

    # Move the bottom constraint Â¬(â„“2 < â„“3) slightly upward
    if constraints_perm[i] == r"$x_2>x_3$":
        ym_off -= 0.05
        xm_off -= .02

    ax4.text(
        xm_off,
        ym_off,
        constraints_perm[i],
        ha="center",
        va="center",
        color=MAIN_BLUE,
        fontsize=fsize
    )

ax4.set_aspect('equal', 'box')
ax4.set_xlim(-1.75, 1.75)
ax4.set_ylim(-1.55, 1.55)
ax4.axis("off")

# ax4.set_title(
#     r"sort input, $L =[\ell_1 = 5,\ell_2 = 4,\ell_3 = 2]$, output = $[2, 4, 5]$",
#     color=HIGHLIGHT_COLOR,
#     fontsize=fsize
# )

plt.tight_layout()
plt.savefig("permutohedron_bubble_style.png", dpi=300, bbox_inches="tight")
plt.savefig("permutohedron_bubble_style.pdf", dpi=300, bbox_inches="tight")
plt.savefig("permutohedron_bubble_style.eps", dpi=300, bbox_inches="tight")
plt.show()

from PyPDF2 import PdfReader, PdfWriter

input_pdf  = "permutohedron_bubble_style.pdf"
output_pdf = "permutohedron_bubble_style.pdf"

# Target size constraints (AMM or other)
target_width_inch  = 5.0   # max width in inches
target_height_inch = 8.0   # max height in inches

# Convert inches to PDF points (1 in = 72 pt)
target_w_pt = target_width_inch  * 72
target_h_pt = target_height_inch * 72

reader = PdfReader(input_pdf)
writer = PdfWriter()
page   = reader.pages[0]

# Original page size in points
orig_w_pt = float(page.mediabox.width)
orig_h_pt = float(page.mediabox.height)

# Compute uniform scale to fit within box
scale_w = target_w_pt / orig_w_pt
scale_h = target_h_pt / orig_h_pt
scale   = min(scale_w, scale_h)

# Scale page vectorially
page.scale_by(scale)

# Write out the new PDF
writer.add_page(page)
with open(output_pdf, "wb") as f:
    writer.write(f)

print(f"Saved '{output_pdf}' scaled to fit within {target_width_inch}\"Ã—{target_height_inch}\" inches (vector-only).")


# %%
# !pip install matplotlib==3.3.4  # Uncomment if needed

import matplotlib.pyplot as plt
from itertools import permutations
import numpy as np
from mpl_toolkits.mplot3d import proj3d
from matplotlib import rcParams
from matplotlib.patches import FancyArrowPatch

fsize = 18

rcParams.update({
    'font.family': 'Times New Roman',
    'font.size': 14,
    'axes.titlesize': 14,
    'axes.labelsize': 14,
    'xtick.labelsize': 14,
    'ytick.labelsize': 14,
    'legend.fontsize': 14,
    'lines.linewidth': 1
})

# ===== COLOR SCHEME =====
MAIN_BLUE = '#4A6E8A'
SECONDARY_BLUE = '#A4B8C9'
HIGHLIGHT_COLOR = '#2C3E50'
GRAY_LIGHT = '#D3D3D3'
GRAY_MEDIUM = '#A9A9A9'
GRAY_DARK = '#5A5A5A'
BLACK = '#000000'
WHITE = '#FFFFFF'

# ===== FONT SIZE =====
MAIN_FONT_SIZE = 22
TITLE_FONT_SIZE = MAIN_FONT_SIZE + 2
SMALL_FONT_SIZE = MAIN_FONT_SIZE - 4

# --- Helper: 3D arrow class ---
class Arrow3D(FancyArrowPatch):
    def __init__(self, xs, ys, zs, *args, **kwargs):
        super().__init__((0, 0), (0, 0), *args, **kwargs)
        self._verts3d = xs, ys, zs

    def draw(self, renderer):
        xs3d, ys3d, zs3d = self._verts3d
        xs, ys, zs = proj3d.proj_transform(xs3d, ys3d, zs3d, self.axes.M)
        self.set_positions((xs[0], ys[0]), (xs[1], ys[1]))
        super().draw(renderer)

    def do_3d_projection(self, renderer=None):
        xs3d, ys3d, zs3d = self._verts3d
        xs, ys, zs = proj3d.proj_transform(xs3d, ys3d, zs3d, self.axes.M)
        self.set_positions((xs[0], ys[0]), (xs[1], ys[1]))
        return np.min(zs)

# Generate all permutations of (1, 2, 3)
perm = list(permutations([1, 2, 3]))
vertices = [list(p) for p in perm]
perm_to_idx = {perm[i]: i for i in range(len(perm))}

# Connect vertices that differ by one adjacent transposition
edges = []
for i, p1 in enumerate(perm):
    for j, p2 in enumerate(perm):
        diff_count = sum(1 for a, b in zip(p1, p2) if a != b)
        if diff_count == 2:
            for k in range(len(p1) - 1):
                if {p1[k], p1[k+1]} == {p2[k], p2[k+1]}:
                    edges.append((i, j))
                    break

# === FIGURE: Initial State with Emphasized Line ===

fig = plt.figure(figsize=(7, 7))
ax = fig.add_subplot(111, projection='3d')

# --- Initial state: vertices and edges (UPDATED VERTEX STYLE ONLY) ---
for i, v in enumerate(vertices):
    ax.scatter(
        v[0], v[1], v[2],
        s=122,
        facecolor=SECONDARY_BLUE,
        edgecolor=BLACK,
        linewidth=1.0
    )

    #print(type(v[0]))
    #print(v)
    if v[0] == 3 and v[1] == 1 and v[2] == 2:
        ax.text(
            v[0] - .5,
            v[1] - .2,
            v[2],
            str(tuple(v)),
            fontsize=MAIN_FONT_SIZE
        )
    else:
        ax.text(
            v[0] + .1,
            v[1],
            v[2] + .05,
            str(tuple(v)),
            fontsize=MAIN_FONT_SIZE
        )

# base edges
for i, j in edges:
    ax.plot(
        [vertices[i][0], vertices[j][0]],
        [vertices[i][1], vertices[j][1]],
        [vertices[i][2], vertices[j][2]],
        color=GRAY_MEDIUM,
        alpha=0.7,
        linewidth=3
    )

# --- Highlight the edge traversal  (3,2,1) â†’ (2,3,1) â†’ (2,1,3) â†’ (1,2,3) ---
path_perm = [(3, 2, 1), (2, 3, 1), (2, 1, 3), (1, 2, 3)]
edge_labels = [
    r"$x_1>x_2$",
    r"$x_2>x_3$",
    r"$x_1>x_2$",
]

# offsets you can tune
dx_top = -1      # move top Â¬(l1<l2) left (currently unused but kept)
dz_bottom = 0.25 # move bottom Â¬(l1<l2) up (currently unused but kept)

for k in range(len(path_perm) - 1):
    i = perm_to_idx[path_perm[k]]
    j = perm_to_idx[path_perm[k + 1]]
    v_i = np.array(vertices[i], dtype=float)
    v_j = np.array(vertices[j], dtype=float)

    # highlighted edge
    ax.plot(
        [v_i[0], v_j[0]],
        [v_i[1], v_j[1]],
        [v_i[2], v_j[2]],
        color=HIGHLIGHT_COLOR,
        linewidth=5.0,
    )

    # --- Edge label (same as 2D hex figure) ---
    mid = 0.5 * (v_i + v_j)
    center = np.array([2.0, 2.0, 2.0])
    direction = mid - center
    # small radial push so text is off the edge
    mid_off = mid + 0.18 * direction

    # k = 0 : bottom Â¬(l1<l2)   â†’ move up in z
    if k == 0:
        mid_off[0] += -.14  # horizontal
        mid_off[2] += .25  # vertical

    if k == 1:
        mid_off[0] += -.7
        mid_off[2] += -.55

    # k = 2 : top Â¬(l1<l2)      â†’ move left in x
    if k == 2:
        mid_off[0] += -.7
        mid_off[2] += -.5

    ax.text(
        mid_off[0],
        mid_off[1],
        mid_off[2],
        edge_labels[k],
        color=BLACK,
        fontsize=MAIN_FONT_SIZE
    )

# # --- Emphasized arrow from (3,2,1) to (1,2,3) (unchanged, still optional) ---
# start = np.array([3, 2, 1])
# end = np.array([1, 2, 3])
#
# arrow = Arrow3D(
#     [start[0], end[0]],
#     [start[1], end[1]],
#     [start[2], end[2]],
#     mutation_scale=20,      # size of the arrow head
#     lw=3,                   # thicker than other edges
#     arrowstyle='-|>',
#     color=HIGHLIGHT_COLOR
# )
# ax.add_artist(arrow)

# Force axes to show [1,2,3] only
ax.set_xticks([1, 2, 3])
ax.set_yticks([1, 2, 3])
ax.set_zticks([1, 2, 3])
ax.set_xlim([1, 3])
ax.set_ylim([1, 3])
ax.set_zlim([1, 3])
ax.tick_params(axis='x', labelsize=16)
ax.tick_params(axis='y', labelsize=16)
ax.tick_params(axis='z', labelsize=16)

plt.tight_layout()
plt.savefig('permutohedron_initial_with_walk.png', dpi=300, bbox_inches='tight')
plt.savefig('permutohedron_initial_with_walk.pdf', dpi=300, bbox_inches='tight')
plt.savefig('permutohedron_initial_with_walk.eps', dpi=300, bbox_inches='tight')

plt.show()


from PyPDF2 import PdfReader, PdfWriter

input_pdf  = "permutohedron_initial_with_walk.pdf"
output_pdf = "permutohedron_initial_with_walk.pdf"

# Target size constraints (AMM or other)
target_width_inch  = 5.0   # max width in inches
target_height_inch = 8.0   # max height in inches

# Convert inches to PDF points (1 in = 72 pt)
target_w_pt = target_width_inch  * 72
target_h_pt = target_height_inch * 72

reader = PdfReader(input_pdf)
writer = PdfWriter()
page   = reader.pages[0]

# Original page size in points
orig_w_pt = float(page.mediabox.width)
orig_h_pt = float(page.mediabox.height)

# Compute uniform scale to fit within box
scale_w = target_w_pt / orig_w_pt
scale_h = target_h_pt / orig_h_pt
scale   = min(scale_w, scale_h)

# Scale page vectorially
page.scale_by(scale)

# Write out the new PDF
writer.add_page(page)
with open(output_pdf, "wb") as f:
    writer.write(f)

print(f"Saved '{output_pdf}' scaled to fit within {target_width_inch}\"Ã—{target_height_inch}\" inches (vector-only).")



# %%
# !pip install matplotlib==3.3.4  # Uncomment if needed

import matplotlib.pyplot as plt
from itertools import permutations
import numpy as np
from mpl_toolkits.mplot3d import proj3d
from matplotlib import rcParams
# from matplotlib.ticker import ScalarFormatter

rcParams.update({
    'font.family': 'Times New Roman',
    'font.size': 14,
    'axes.titlesize': 14,
    'axes.labelsize': 14,
    'xtick.labelsize': 14,
    'ytick.labelsize': 14,
    'legend.fontsize': 14,
    'lines.linewidth': 1
})

# ===== COLOR SCHEME =====
MAIN_BLUE = '#4A6E8A'
SECONDARY_BLUE = '#A4B8C9'
HIGHLIGHT_COLOR = '#2C3E50'
GRAY_LIGHT = '#D3D3D3'
GRAY_MEDIUM = '#A9A9A9'
GRAY_DARK = '#5A5A5A'
BLACK = '#000000'
WHITE = '#FFFFFF'

# ===== FONT SIZE =====
MAIN_FONT_SIZE = 22
TITLE_FONT_SIZE = MAIN_FONT_SIZE + 2
SMALL_FONT_SIZE = MAIN_FONT_SIZE - 4

# ===== NODE STYLE =====
ACTIVE_NODE_SIZE = 140      # larger, emphasized nodes
INACTIVE_NODE_SIZE = 90     # smaller, faded nodes
NODE_EDGE_WIDTH = 1.2

# Generate all permutations of (1, 2, 3)
perm = list(permutations([1, 2, 3]))
vertices = [list(p) for p in perm]
perm_to_idx = {perm[i]: i for i in range(len(perm))}

# Connect vertices that differ by one adjacent transposition
edges = []
for i, p1 in enumerate(perm):
    for j, p2 in enumerate(perm):
        diff_count = sum(1 for a, b in zip(p1, p2) if a != b)
        if diff_count == 2:
            for k in range(len(p1) - 1):
                if {p1[k], p1[k+1]} == {p2[k], p2[k+1]}:
                    edges.append((i, j))
                    break


def draw_vertex_label(ax, v):
    text_kwargs = {}
    if v[0] == 3 and v[1] == 1 and v[2] == 2:
        pos = (v[0] - .65, v[1] - .2, v[2])
    elif v[0] == 2 and v[1] == 1 and v[2] == 3:
        pos = (v[0] - .65, v[1] - .2, v[2])
    elif v[0] == 3 and v[1] == 2 and v[2] == 1:
        pos = (v[0] + .02, v[1] - .36, v[2] + .02)
        text_kwargs = {"ha": "right", "va": "center"}
    else:
        pos = (v[0] + .1, v[1], v[2] + .05)

    ax.text(*pos, str(tuple(v)), fontsize=MAIN_FONT_SIZE, color=BLACK,
            zorder=20, clip_on=False, **text_kwargs)


# === FIGURE: Sequence of Constraint Pruning ===

fig2 = plt.figure(figsize=(15, 15))

ax3 = fig2.add_subplot(131, projection='3d')
ax4 = fig2.add_subplot(132, projection='3d')
ax5 = fig2.add_subplot(133, projection='3d')

# === Left: Initial State ===
for i, v in enumerate(vertices):
    ax3.scatter(
        *v,
        s=ACTIVE_NODE_SIZE,
        facecolor=SECONDARY_BLUE,
        edgecolor=BLACK,
        linewidth=NODE_EDGE_WIDTH
    )

for i, j in edges:
    ax3.plot(
        [vertices[i][0], vertices[j][0]],
        [vertices[i][1], vertices[j][1]],
        [vertices[i][2], vertices[j][2]],
        color=GRAY_MEDIUM,
        alpha=0.7
    )

for v in vertices:
    draw_vertex_label(ax3, v)

ax3.set_title("Initial State\n(6 rank words)", fontsize=TITLE_FONT_SIZE)
ax3.set_xlabel(r"$x_1$", fontsize=SMALL_FONT_SIZE)
ax3.set_ylabel(r"$x_2$", fontsize=SMALL_FONT_SIZE)
# ax3.set_zlabel(r"$x_3$", fontsize=SMALL_FONT_SIZE)

# === Middle: After xâ‚ < xâ‚‚ ===
valid_indices = [i for i, p in enumerate(perm) if p[0] < p[1]]
for i, v in enumerate(vertices):
    valid = i in valid_indices
    ax4.scatter(
        *v,
        s=ACTIVE_NODE_SIZE if valid else INACTIVE_NODE_SIZE,
        facecolor=SECONDARY_BLUE if valid else GRAY_LIGHT,
        edgecolor=BLACK if valid else GRAY_MEDIUM,
        linewidth=NODE_EDGE_WIDTH if valid else 0.5,
        alpha=1.0 if valid else 0.35
    )

for i, j in edges:
    color = GRAY_DARK if i in valid_indices and j in valid_indices else GRAY_MEDIUM
    alpha = 0.7 if i in valid_indices and j in valid_indices else 0.2
    ax4.plot(
        [vertices[i][0], vertices[j][0]],
        [vertices[i][1], vertices[j][1]],
        [vertices[i][2], vertices[j][2]],
        color=color,
        alpha=alpha
    )

# Plane x1 = x2 (filled)
xx, zz = np.meshgrid(range(0, 4), range(0, 4))
yy = xx
ax4.plot_surface(xx, yy, zz, alpha=0.07, color=SECONDARY_BLUE)

# Dashed border for x1 = x2 plane
ax4.plot_wireframe(
    xx, yy, zz,
    color=GRAY_DARK,
    linewidth=1.0,
    linestyle='--',
    alpha=0.55
)

for v in vertices:
    draw_vertex_label(ax4, v)

ax4.set_title(r"After $x_1 < x_2$ Constraint" + "\n(3 rank words left)",
              fontsize=TITLE_FONT_SIZE)
ax4.set_xlabel(r"$x_1$", fontsize=SMALL_FONT_SIZE)
ax4.set_ylabel(r"$x_2$", fontsize=SMALL_FONT_SIZE)
# ax4.set_zlabel(r"$x_3$", fontsize=SMALL_FONT_SIZE)

# === Right: After xâ‚‚ < xâ‚ƒ ===
valid_indices_final = [i for i, p in enumerate(perm) if p[0] < p[1] and p[1] < p[2]]
for i, v in enumerate(vertices):
    valid = i in valid_indices_final
    ax5.scatter(
        *v,
        s=ACTIVE_NODE_SIZE if valid else INACTIVE_NODE_SIZE,
        facecolor=HIGHLIGHT_COLOR if valid else GRAY_LIGHT,
        edgecolor=BLACK if valid else GRAY_MEDIUM,
        linewidth=NODE_EDGE_WIDTH if valid else 0.5,
        alpha=1.0 if valid else 0.35
    )

for i, j in edges:
    color = HIGHLIGHT_COLOR if i in valid_indices_final and j in valid_indices_final else GRAY_MEDIUM
    alpha = 0.7 if i in valid_indices_final and j in valid_indices_final else 0.2
    ax5.plot(
        [vertices[i][0], vertices[j][0]],
        [vertices[i][1], vertices[j][1]],
        [vertices[i][2], vertices[j][2]],
        color=color,
        alpha=alpha
    )

# Plane x1 = x2 (filled)
xx, zz = np.meshgrid(range(0, 4), range(0, 4))
yy = xx
ax5.plot_surface(xx, yy, zz, alpha=0.07, color=SECONDARY_BLUE)

# Dashed border for x1 = x2 plane
ax5.plot_wireframe(
    xx, yy, zz,
    color=GRAY_DARK,
    linewidth=1.0,
    linestyle='--',
    alpha=0.55
)

# Plane x2 = x3 (filled)
xx, yy = np.meshgrid(range(0, 4), range(0, 4))
zz = yy
ax5.plot_surface(xx, yy, zz, alpha=0.07, color=GRAY_MEDIUM)

# Dashed border for x2 = x3 plane
ax5.plot_wireframe(
    xx, yy, zz,
    color=GRAY_DARK,
    linewidth=1.0,
    linestyle='--',
    alpha=0.55
)

for v in vertices:
    draw_vertex_label(ax5, v)

ax5.set_title(r"After $x_2 < x_3$ Constraint" + "\n(Only sorted rank word left)",
              fontsize=TITLE_FONT_SIZE)
ax5.set_xlabel(r"$x_1$", fontsize=SMALL_FONT_SIZE)
ax5.set_ylabel(r"$x_2$", fontsize=SMALL_FONT_SIZE)
# ax5.set_zlabel(r"$x_3$", fontsize=SMALL_FONT_SIZE)

# Force axes to show [1,2,3] only and soften grid/panes
for ax in [ax3, ax4, ax5]:
    ax.set_xticks([1, 2, 3])
    ax.set_yticks([1, 2, 3])
    ax.set_zticks([1, 2, 3])
    ax.set_xlim([1, 3])
    ax.set_ylim([1, 3])
    ax.set_zlim([1, 3])
    ax.tick_params(axis='x', labelsize=19)
    ax.tick_params(axis='y', labelsize=19)
    ax.tick_params(axis='z', labelsize=19)

    # softer grid
    ax.grid(True)
    ax.xaxis._axinfo["grid"]['color'] = (0.85, 0.85, 0.85, 0.30)
    ax.yaxis._axinfo["grid"]['color'] = (0.85, 0.85, 0.85, 0.30)
    ax.zaxis._axinfo["grid"]['color'] = (0.85, 0.85, 0.85, 0.30)

    # transparent panes
    ax.xaxis.set_pane_color((1, 1, 1, 0.0))
    ax.yaxis.set_pane_color((1, 1, 1, 0.0))
    ax.zaxis.set_pane_color((1, 1, 1, 0.0))

plt.tight_layout()
plt.savefig('permutohedron_sequence.png', dpi=300, bbox_inches='tight')
plt.savefig("permutohedron_sequence.pdf", dpi=300, bbox_inches='tight')
plt.savefig("permutohedron_sequence.eps", dpi=300, bbox_inches='tight')


# %%
from PyPDF2 import PdfReader, PdfWriter

input_pdf  = "permutohedron_sequence.pdf"
output_pdf = "permutohedron_sequence_scaled.pdf"

# Target size constraints (AMM or other)
target_width_inch  = 5.0   # max width in inches
target_height_inch = 8.0   # max height in inches

# Convert inches to PDF points (1 in = 72 pt)
target_w_pt = target_width_inch  * 72
target_h_pt = target_height_inch * 72

reader = PdfReader(input_pdf)
writer = PdfWriter()
page   = reader.pages[0]

# Original page size in points
orig_w_pt = float(page.mediabox.width)
orig_h_pt = float(page.mediabox.height)

# Compute uniform scale to fit within box
scale_w = target_w_pt / orig_w_pt
scale_h = target_h_pt / orig_h_pt
scale   = min(scale_w, scale_h)

# Scale page vectorially
page.scale_by(scale)

# Write out the new PDF
writer.add_page(page)
with open(output_pdf, "wb") as f:
    writer.write(f)

print(f"Saved '{output_pdf}' scaled to fit within {target_width_inch}\"Ã—{target_height_inch}\" inches (vector-only).")


# %% [markdown]
# #### gradient flow on the permutohedron

# %%
# !pip install matplotlib==3.3.4  # Uncomment if needed

import matplotlib.pyplot as plt
from itertools import permutations
import numpy as np
from mpl_toolkits.mplot3d import proj3d
from matplotlib import rcParams
from matplotlib.patches import FancyArrowPatch

rcParams.update({
    'font.family': 'Times New Roman',
    'font.size': 18,
    'axes.titlesize': 18,
    'axes.labelsize': 18,
    'xtick.labelsize': 18,
    'ytick.labelsize': 18,
    'legend.fontsize': 18,
    'lines.linewidth': 1
})

# ===== COLOR SCHEME =====
MAIN_BLUE = '#4A6E8A'
SECONDARY_BLUE = '#A4B8C9'
HIGHLIGHT_COLOR = '#2C3E50'
GRAY_LIGHT = '#D3D3D3'
GRAY_MEDIUM = '#A9A9A9'
GRAY_DARK = '#5A5A5A'
BLACK = '#000000'
WHITE = '#FFFFFF'

# ===== FONT SIZE =====
MAIN_FONT_SIZE = 22
TITLE_FONT_SIZE = MAIN_FONT_SIZE + 2
SMALL_FONT_SIZE = MAIN_FONT_SIZE - 4

# --- Helper: 3D arrow class ---
class Arrow3D(FancyArrowPatch):
    def __init__(self, xs, ys, zs, *args, **kwargs):
        super().__init__((0, 0), (0, 0), *args, **kwargs)
        self._verts3d = xs, ys, zs

    def draw(self, renderer):
        xs3d, ys3d, zs3d = self._verts3d
        xs, ys, zs = proj3d.proj_transform(xs3d, ys3d, zs3d, self.axes.M)
        self.set_positions((xs[0], ys[0]), (xs[1], ys[1]))
        super().draw(renderer)

    def do_3d_projection(self, renderer=None):
        xs3d, ys3d, zs3d = self._verts3d
        xs, ys, zs = proj3d.proj_transform(xs3d, ys3d, zs3d, self.axes.M)
        self.set_positions((xs[0], ys[0]), (xs[1], ys[1]))
        return np.min(zs)

# Generate all permutations of (1, 2, 3)
perm = list(permutations([1, 2, 3]))
vertices = [list(p) for p in perm]
perm_to_idx = {perm[i]: i for i in range(len(perm))}

# Connect vertices that differ by one adjacent transposition
edges = []
for i, p1 in enumerate(perm):
    for j, p2 in enumerate(perm):
        diff_count = sum(1 for a, b in zip(p1, p2) if a != b)
        if diff_count == 2:
            for k in range(len(p1) - 1):
                if {p1[k], p1[k+1]} == {p2[k], p2[k+1]}:
                    edges.append((i, j))
                    break

# === FIGURE: Initial State with Emphasized Line ===

fig = plt.figure(figsize=(7, 7))
ax = fig.add_subplot(111, projection='3d')

# --- Initial state: vertices and edges (UPDATED VERTEX STYLE) ---
for i, v in enumerate(vertices):
    # node circle style consistent with reference code
    ax.scatter(
        v[0], v[1], v[2],
        s=122,
        facecolor=SECONDARY_BLUE,
        edgecolor=BLACK,
        linewidth=1.0
    )

    # label shifting consistent with reference code
    if v[0] == 3 and v[1] == 1 and v[2] == 2:
        ax.text(
            v[0] - .5,
            v[1] - .2,
            v[2],
            str(tuple(v)),
            fontsize=MAIN_FONT_SIZE
        )
    elif v[0] == 2 and v[1] == 1 and v[2] == 3:
        ax.text(
            v[0] - .5,
            v[1] - .2,
            v[2],
            str(tuple(v)),
            fontsize=MAIN_FONT_SIZE
        )
    else:
        ax.text(
            v[0] + .1,
            v[1],
            v[2] + .05,
            str(tuple(v)),
            fontsize=MAIN_FONT_SIZE
        )

# base edges
for i, j in edges:
    ax.plot(
        [vertices[i][0], vertices[j][0]],
        [vertices[i][1], vertices[j][1]],
        [vertices[i][2], vertices[j][2]],
        color=GRAY_MEDIUM,
        alpha=0.7,
        linewidth=3  # to match the reference style
    )

# --- Highlight the edge traversal  (3,2,1) â†’ (2,3,1) â†’ (2,1,3) â†’ (1,2,3) ---
path_perm = [(3, 2, 1), (2, 3, 1), (2, 1, 3), (1, 2, 3)]

for k in range(len(path_perm) - 1):
    i = perm_to_idx[path_perm[k]]
    j = perm_to_idx[path_perm[k + 1]]
    ax.plot(
        [vertices[i][0], vertices[j][0]],
        [vertices[i][1], vertices[j][1]],
        [vertices[i][2], vertices[j][2]],
        color=HIGHLIGHT_COLOR,
        linewidth=7,   # thicker path
        alpha=.5
    )

# --- Emphasized arrow from (3,2,1) to (1,2,3) ---
start = np.array([3, 2, 1])
end = np.array([1, 2, 3])

arrow = Arrow3D(
    [start[0], end[0]],
    [start[1], end[1]],
    [start[2], end[2]],
    mutation_scale=30,      # size of the arrow head
    lw=7,                   # thicker than other edges
    arrowstyle='-|>',
    color=HIGHLIGHT_COLOR
)
ax.add_artist(arrow)

# Force axes to show [1,2,3] only
ax.set_xticks([1, 2, 3])
ax.set_yticks([1, 2, 3])
ax.set_zticks([1, 2, 3])
ax.set_xlim([1, 3])
ax.set_ylim([1, 3])
ax.set_zlim([1, 3])
ax.tick_params(axis='x', labelsize=16)
ax.tick_params(axis='y', labelsize=16)
ax.tick_params(axis='z', labelsize=16)

plt.tight_layout()
plt.savefig('gradient_flow_straight_line.png', dpi=300, bbox_inches='tight')
plt.savefig('gradient_flow_straight_line.pdf', dpi=300, bbox_inches='tight')
plt.savefig('gradient_flow_straight_line.eps', dpi=300, bbox_inches='tight')

#plt.show()

from PyPDF2 import PdfReader, PdfWriter

input_pdf  = "gradient_flow_straight_line.pdf"
output_pdf = "gradient_flow_straight_line.pdf"

# Target size constraints (AMM or other)
target_width_inch  = 5.0   # max width in inches
target_height_inch = 8.0   # max height in inches

# Convert inches to PDF points (1 in = 72 pt)
target_w_pt = target_width_inch  * 72
target_h_pt = target_height_inch * 72

reader = PdfReader(input_pdf)
writer = PdfWriter()
page   = reader.pages[0]

# Original page size in points
orig_w_pt = float(page.mediabox.width)
orig_h_pt = float(page.mediabox.height)

# Compute uniform scale to fit within box
scale_w = target_w_pt / orig_w_pt
scale_h = target_h_pt / orig_h_pt
scale   = min(scale_w, scale_h)

# Scale page vectorially
page.scale_by(scale)

# Write out the new PDF
writer.add_page(page)
with open(output_pdf, "wb") as f:
    writer.write(f)

print(
    f"Saved '{output_pdf}' scaled to fit within "
    f"{target_width_inch}\"Ã—{target_height_inch}\" inches (vector-only)."
)


# %%
import numpy as np
import matplotlib.pyplot as plt
from mpl_toolkits.mplot3d import Axes3D
from matplotlib import rcParams

# --- Start & End node color (slightly darker, navy-leaning) ---
DARK_NODE = "#375A7F"     # dark blue, close to navy

rcParams.update({
    'font.family': 'Times New Roman',
    'font.size': 18,
    'axes.titlesize': 18,
    'axes.labelsize': 18,
    'xtick.labelsize': 18,
    'ytick.labelsize': 18,
    'legend.fontsize': 18,
    'lines.linewidth': 1
})


# Parameters
n = 3
R0 = 8.0
r0 = np.sqrt(R0)
V0 = 0.5 * r0**2
eps = 0.01
T_eps = np.log(r0 / eps)
t_max = np.ceil(T_eps)
nlogn = n * np.log(n)

def r_flow(t, r0):
    t = np.asarray(t)
    return r0 * np.exp(-t)

def V_of_r(r):
    r = np.asarray(r)
    return 0.5 * r**2

# Data generation
t_dense = np.linspace(0.0, T_eps, 300)
r_dense = r_flow(t_dense, r0)
V_dense = V_of_r(r_dense)

k_steps = np.arange(0, int(t_max) + 1)
r_steps = r_flow(k_steps, r0)
V_steps = V_of_r(r_steps)

r_vals = np.linspace(0.0, r0, 100)
t_vals = np.linspace(0.0, t_max, 100)
R_grid, T_grid = np.meshgrid(r_vals, t_vals)
V_grid = V_of_r(R_grid)

# Plotting - INCREASED figure width to prevent cutoff
fig = plt.figure(figsize=(9, 9))  # Changed from (7, 7) to (9, 7)
ax = fig.add_subplot(111, projection='3d')

ax.plot_surface(R_grid, T_grid, V_grid, alpha=0.3)
ax.plot(r_dense, t_dense, V_dense, linewidth=3, alpha=0.8, color="black")

# --- UPDATED NODE STYLE (matches permutohedron nodes) ---
ax.scatter(
    r_steps[:-1],
    k_steps[:-1],
    V_steps[:-1],
    s=122,
    facecolor="navy",   # SECONDARY_BLUE
    edgecolor="#000000",   # BLACK
    linewidth=1.5
)

# Markers
start_t = 0.0
start_r = r0
start_V = V0
stop_t = T_eps
stop_r = r_flow(stop_t, r0)
stop_V = V_of_r(stop_r)

ax.scatter(
    [start_r],
    [start_t],
    [start_V],
    s=122,
    facecolor="black",
    edgecolor="#000000",
    linewidth=1.5
)
ax.scatter(
    [stop_r],
    [stop_t],
    [stop_V],
    s=122,
    facecolor="black",
    edgecolor="#000000",
    linewidth=1.5
)

ax.text(
    start_r - .12,
    start_t,
    start_V - .25,
    rf"($r={start_r:.2f},$" + "\n" + r"    $t=0$)",
    ha="left",
    va="top",
    fontsize=25
)
ax.text(
    stop_r,
    stop_t,
    stop_V + 0.2,
    r"$(r \approx \varepsilon, t = T_\varepsilon)$",
    ha="center",
    va="bottom",
    fontsize=25
)

# Labels
ax.set_xlabel(r"$r=\|x - v_s\|_2$", fontsize=25, labelpad = 10)
ax.set_ylabel(r"$t$", fontsize=25)
# Remove default z-label
ax.set_zlabel("")

# Compute the top of the z-axis in 3D coords
z_top = ax.get_zlim()[1]
x0, y0 = ax.get_xlim()[0], ax.get_ylim()[0]

# Add label above the z axis
ax.text(
    x0 + 4.7,                     # anchor x: left edge of plot
    y0,                           # anchor y: bottom edge of plot
    z_top + 0.1 * (z_top) + 3.2,  # slightly above top of z-axis
    r"$V(r(t))$",
    ha="center",
    va="bottom",
    fontsize=25
)

ax.set_yticks([0, t_max])
ax.set_yticklabels([r"0", r"$T_\varepsilon$"])

# Saving
plt.savefig('final_gradient_path.png', dpi=300, bbox_inches='tight')
plt.savefig('final_gradient_path.pdf', dpi=300, bbox_inches='tight')
plt.savefig('final_gradient_path.eps', dpi=300, bbox_inches='tight')

# plt.show()

from PyPDF2 import PdfReader, PdfWriter

input_pdf  = "final_gradient_path.pdf"
output_pdf = "final_gradient_path.pdf"

# Target size constraints (AMM or other)
target_width_inch  = 5.0   # max width in inches
target_height_inch = 8.0   # max height in inches

# Convert inches to PDF points (1 in = 72 pt)
target_w_pt = target_width_inch  * 72
target_h_pt = target_height_inch * 72

reader = PdfReader(input_pdf)
writer = PdfWriter()
page   = reader.pages[0]

# Original page size in points
orig_w_pt = float(page.mediabox.width)
orig_h_pt = float(page.mediabox.height)

# Compute uniform scale to fit within box
scale_w = target_w_pt / orig_w_pt
scale_h = target_h_pt / orig_h_pt
scale   = min(scale_w, scale_h)

# Scale page vectorially
page.scale_by(scale)

# Write out the new PDF
writer.add_page(page)
with open(output_pdf, "wb") as f:
    writer.write(f)

print(
    f"Saved '{output_pdf}' scaled to fit within "
    f"{target_width_inch}\"Ã—{target_height_inch}\" inches (vector-only)."
)


# %% [markdown]
# # DEBUG - tangent cone

# %%
# !pip install matplotlib==3.3.4  # Uncomment if needed

import matplotlib.pyplot as plt
import numpy as np
from mpl_toolkits.mplot3d import proj3d
from matplotlib import rcParams
from matplotlib.patches import FancyArrowPatch
from PyPDF2 import PdfReader, PdfWriter

# ===== STYLE SETTINGS =====
rcParams.update({
    'font.family': 'Times New Roman',
    'font.size': 18,
    'axes.titlesize': 18,
    'axes.labelsize': 18,
    'xtick.labelsize': 18,
    'ytick.labelsize': 18,
    'legend.fontsize': 18,
    'lines.linewidth': 1
})

# ===== COLOR SCHEME =====
MAIN_BLUE      = '#4A6E8A'
SECONDARY_BLUE = '#A4B8C9'
HIGHLIGHT_COLOR = '#2C3E50'
GRAY_LIGHT     = '#D3D3D3'
GRAY_MEDIUM    = '#A9A9A9'
GRAY_DARK      = '#5A5A5A'
BLACK          = '#000000'
WHITE          = '#FFFFFF'

MAIN_FONT_SIZE   = 22
TITLE_FONT_SIZE  = MAIN_FONT_SIZE + 2
SMALL_FONT_SIZE  = MAIN_FONT_SIZE - 4
LABEL_FONT_SIZE  = 18   # one consistent label size

# --- Helper: 3D arrow class ---------------------------------------------------
class Arrow3D(FancyArrowPatch):
    def __init__(self, xs, ys, zs, *args, **kwargs):
        super().__init__((0, 0), (0, 0), *args, **kwargs)
        self._verts3d = xs, ys, zs

    def draw(self, renderer):
        xs3d, ys3d, zs3d = self._verts3d
        xs, ys, zs = proj3d.proj_transform(xs3d, ys3d, zs3d, self.axes.M)
        self.set_positions((xs[0], ys[0]), (xs[1], ys[1]))
        super().draw(renderer)

    def do_3d_projection(self, renderer=None):
        xs3d, ys3d, zs3d = self._verts3d
        xs, ys, zs = proj3d.proj_transform(xs3d, ys3d, zs3d, self.axes.M)
        self.set_positions((xs[0], ys[0]), (xs[1], ys[1]))
        return np.min(zs)

# === GEOMETRY: boundary point and gradient ====================================

# Boundary state x(t) on the facet x1 = x2
origin = np.array([2.0, 2.0, 2.0])

# Normal to plane x1 = x2
normal_vec = np.array([1.0, -1.0, 0.0])
normal_vec /= np.linalg.norm(normal_vec)

# Unconstrained -âˆ‡V direction (has both tangential + normal components)
unconstrained_vec = np.array([1.3, -0.4, 1.0])
unconstrained_vec = unconstrained_vec / np.linalg.norm(unconstrained_vec) * 1.4

# Decompose into tangential + normal
dot_prod         = np.dot(unconstrained_vec, normal_vec)
normal_component = dot_prod * normal_vec
projected_vec    = unconstrained_vec - normal_component

end_unc  = origin + unconstrained_vec   # tip of unconstrained gradient
end_proj = origin + projected_vec       # tip of projected flow
diss_vec = end_unc - end_proj           # the "lost" / dissipated component

# === PLANE PATCH (x1 = x2), SHRUNK AROUND ORIGIN =============================

plane_half = 1  # controls size of plane patch

x_vals = np.linspace(origin[0] - plane_half,
                     origin[0] + plane_half, 4)
z_vals = np.linspace(origin[2] - plane_half,
                     origin[2] + plane_half, 4)

xx, zz = np.meshgrid(x_vals, z_vals)
yy = xx  # still x1 = x2 plane

# === PLOT =====================================================================

fig = plt.figure(figsize=(7, 7))
ax = fig.add_subplot(111, projection='3d')

# Plane style: matching your x1 = x2 permutohedron facet
ax.plot_surface(
    xx, yy, zz,
    alpha=0.15,
    color=SECONDARY_BLUE
)
ax.plot_wireframe(
    xx, yy, zz,
    color=GRAY_DARK,
    linewidth=1.0,
    linestyle='--',
    alpha=0.4
)

# --- Key vertices -------------------------------------------------------------

# 1. Boundary state node (on plane)
ax.scatter(
    origin[0], origin[1], origin[2],
    s=122,
    facecolor=SECONDARY_BLUE,
    edgecolor=BLACK,
    linewidth=1.0,
    zorder=6
)

# 2. Projected tip (on plane)
ax.scatter(
    end_proj[0], end_proj[1], end_proj[2],
    s=122,
    facecolor=SECONDARY_BLUE,
    edgecolor=BLACK,
    linewidth=1.0,
    zorder=6
)

# 3. Unconstrained tip (off plane)
ax.scatter(
    end_unc[0], end_unc[1], end_unc[2],
    s=122,
    facecolor=GRAY_LIGHT,
    edgecolor=BLACK,
    linewidth=1.0,
    zorder=6
)

# --- Arrows: unconstrained, projected, dissipation ---------------------------

# Unconstrained -âˆ‡V (black, dashed)
arrow_unc = Arrow3D(
    [origin[0], end_unc[0]],
    [origin[1], end_unc[1]],
    [origin[2], end_unc[2]],
    mutation_scale=24,
    lw=6,
    linestyle='--',
    arrowstyle='-|>',
    color="black",
    alpha=0.6,
)
ax.add_artist(arrow_unc)

# Projected flow (black, solid)
arrow_proj = Arrow3D(
    [origin[0], end_proj[0]],
    [origin[1], end_proj[1]],
    [origin[2], end_proj[2]],
    mutation_scale=30,
    lw=6,
    arrowstyle='-|>',
    color="black",
    alpha=1.0,
)
ax.add_artist(arrow_proj)

# Dissipation (normal component) from projected tip to unconstrained tip
arrow_diss = Arrow3D(
    [end_proj[0], end_unc[0]],
    [end_proj[1], end_unc[1]],
    [end_proj[2], end_unc[2]],
    mutation_scale=16,
    lw=6,
    linestyle=':',
    arrowstyle='-|>',
    color="black",
    alpha=0.25,
)
ax.add_artist(arrow_diss)

# ============================================================================
# VERTEX LABELS (positions + text) â€“ move these independently
# ============================================================================

# Positions for vertex labels
origin_vert_label_pos = origin   + np.array([-0.1, -0.30, -0.70])
proj_vert_label_pos   = end_proj + np.array([ 0.0,  0.20, -0.10])
unc_vert_label_pos    = end_unc  + np.array([ 0.0,  0.25,  0.10])

# Text for vertex labels (your paperâ€™s example)
origin_vert_text = (
    r'$x_0 = (3,2,1)$' + '\n' +
    r'$L = [5,4,2]$'
)

proj_vert_text = r'$x_{\mathrm{proj}}$'

unc_vert_text = (
    r'$v_s = (1,2,3)$' + '\n' +
    r'$L_s = [2,4,5]$'
)

# Draw vertex labels
ax.text(
    origin_vert_label_pos[0], origin_vert_label_pos[1], origin_vert_label_pos[2] + .55,
    origin_vert_text,
    color="black",
    fontsize=LABEL_FONT_SIZE,
    ha='center', va='top'
)

ax.text(
    proj_vert_label_pos[0] + .05, proj_vert_label_pos[1] + .05, proj_vert_label_pos[2],
    proj_vert_text,
    color="black",
    fontsize=LABEL_FONT_SIZE,
    ha='center', va='bottom'
)

ax.text(
    unc_vert_label_pos[0], unc_vert_label_pos[1], unc_vert_label_pos[2] + .08,
    unc_vert_text,
    color="black",
    fontsize=LABEL_FONT_SIZE,
    ha='center', va='bottom'
)

# ============================================================================
# EDGE LABELS (positions + text) â€“ separate from vertex labels
# ============================================================================

# Positions for edge labels (you can move these freely)
unc_edge_label_pos  = (origin + end_unc)/2   + np.array([ 0.25,  0.10,  0.10])
proj_edge_label_pos = (origin + end_proj)/2  + np.array([ 0.40, -0.15,  0.00])
diss_edge_label_pos = (end_proj + end_unc)/2 + np.array([-0.70,  0.10,  0.10])

constraint_pos = np.array([
    origin[0] + plane_half + 0.3,
    origin[1],
    origin[2] - plane_half
])

# Text for edge labels
unc_edge_text = r'$-\nabla V$' + '\n' + r'(unconstrained)'

proj_edge_text = r'$\dot{x} = \Pi_{T_{\mathcal{P}_n}(x)}[-\nabla V]$'

diss_edge_text = (
    r'Dissipation' + '\n' +
    r'normal to facet' + '\n' + r'$x_i = x_j$'
)

# Draw edge labels
ax.text(
    unc_edge_label_pos[0] - .5, unc_edge_label_pos[1] - .5, unc_edge_label_pos[2],
    unc_edge_text,
    color="black",
    fontsize=LABEL_FONT_SIZE,
    ha='center', va='bottom',
    alpha=0.8
)

ax.text(
    proj_edge_label_pos[0] + .4, proj_edge_label_pos[1] + .2, proj_edge_label_pos[2] - .4,
    proj_edge_text,
    color="black",
    fontsize=LABEL_FONT_SIZE,
    ha='center', va='center',
    alpha=1.0
)

ax.text(
    diss_edge_label_pos[0] + .95, diss_edge_label_pos[1] + .9, diss_edge_label_pos[2],
    diss_edge_text,
    color="black",
    fontsize=LABEL_FONT_SIZE,
    ha='center', va='center',
    alpha=0.6
)

# Constraint boundary label (plane)
ax.text(
    constraint_pos[0], constraint_pos[1], constraint_pos[2],
    r'Constraint Boundary' + '\n' + r'$x_i = x_j$',
    color="black",
    fontsize=LABEL_FONT_SIZE,
    ha='center',
    va='center'
)

# === AXES / 3D GRID ===========================================================

ax.set_xticks([0, 1, 2, 3])
ax.set_yticks([0, 1, 2, 3])
ax.set_zticks([0, 1, 2, 3])

# Hide tick labels for a clean look
ax.set_xticklabels([])
ax.set_yticklabels([])
ax.set_zticklabels([])

# Zoomed-in box so dissipation is clearly visible
ax.set_xlim(1.5, 3.3)
ax.set_ylim(1.5, 3.3)
ax.set_zlim(1.5, 3.3)

ax.grid(True)
ax.xaxis._axinfo["grid"]['color'] = (0.85, 0.85, 0.85, 0.45)
ax.yaxis._axinfo["grid"]['color'] = (0.85, 0.85, 0.85, 0.45)
ax.zaxis._axinfo["grid"]['color'] = (0.85, 0.85, 0.85, 0.45)

ax.xaxis.set_pane_color((1, 1, 1, 0.0))
ax.yaxis.set_pane_color((1, 1, 1, 0.0))
ax.zaxis.set_pane_color((1, 1, 1, 0.0))

ax.view_init(elev=18, azim=135)

plt.tight_layout()

# === SAVE =====================================================================

filename_base = 'tangent_cone_projection_vertex_edge_labels_separate'
plt.savefig(f'{filename_base}.png', dpi=300, bbox_inches='tight')
plt.savefig(f'{filename_base}.pdf', dpi=300, bbox_inches='tight')
plt.savefig(f'{filename_base}.eps', dpi=300, bbox_inches='tight')

input_pdf  = f"{filename_base}.pdf"
output_pdf = f"{filename_base}_scaled.pdf"

target_width_inch  = 5.0
target_height_inch = 5.0
target_w_pt = target_width_inch  * 72
target_h_pt = target_height_inch * 72

try:
    reader = PdfReader(input_pdf)
    writer = PdfWriter()
    page   = reader.pages[0]

    orig_w_pt = float(page.mediabox.width)
    orig_h_pt = float(page.mediabox.height)

    scale_w = target_w_pt / orig_w_pt
    scale_h = target_h_pt / orig_h_pt
    scale   = min(scale_w, scale_h)

    page.scale_by(scale)
    writer.add_page(page)
    with open(output_pdf, "wb") as f:
        writer.write(f)
    print(f"Saved '{output_pdf}' at ~{target_width_inch}\"Ã—{target_height_inch}\".")
except Exception as e:
    print("Skipped PDF scaling:", e)


# %%
# !pip install matplotlib==3.3.4  # Uncomment if needed

import matplotlib.pyplot as plt
import numpy as np
from mpl_toolkits.mplot3d import proj3d
from matplotlib import rcParams
from matplotlib.patches import FancyArrowPatch
from PyPDF2 import PdfReader, PdfWriter

# ===== STYLE SETTINGS =====
rcParams.update({
    'font.family': 'Times New Roman',
    'font.size': 18,
    'axes.titlesize': 18,
    'axes.labelsize': 18,
    'xtick.labelsize': 18,
    'ytick.labelsize': 18,
    'legend.fontsize': 18,
    'lines.linewidth': 1
})

# ===== COLOR SCHEME =====
MAIN_BLUE      = '#4A6E8A'
SECONDARY_BLUE = '#A4B8C9'
HIGHLIGHT_COLOR = '#2C3E50'
GRAY_LIGHT     = '#D3D3D3'
GRAY_MEDIUM    = '#A9A9A9'
GRAY_DARK      = '#5A5A5A'
BLACK          = '#000000'
WHITE          = '#FFFFFF'

MAIN_FONT_SIZE   = 22
TITLE_FONT_SIZE  = MAIN_FONT_SIZE + 2
SMALL_FONT_SIZE  = MAIN_FONT_SIZE - 4
LABEL_FONT_SIZE  = 18   # one consistent label size

# --- Helper: 3D arrow class ---------------------------------------------------
class Arrow3D(FancyArrowPatch):
    def __init__(self, xs, ys, zs, *args, **kwargs):
        super().__init__((0, 0), (0, 0), *args, **kwargs)
        self._verts3d = xs, ys, zs

    def draw(self, renderer):
        xs3d, ys3d, zs3d = self._verts3d
        xs, ys, zs = proj3d.proj_transform(xs3d, ys3d, zs3d, self.axes.M)
        self.set_positions((xs[0], ys[0]), (xs[1], ys[1]))
        super().draw(renderer)

    def do_3d_projection(self, renderer=None):
        xs3d, ys3d, zs3d = self._verts3d
        xs, ys, zs = proj3d.proj_transform(xs3d, ys3d, zs3d, self.axes.M)
        self.set_positions((xs[0], ys[0]), (xs[1], ys[1]))
        return np.min(zs)

# === GEOMETRY: local tangent-cone cartoon at the facet between 321 and 231 ===
# This is a schematic local view; coordinates are chosen for clarity, not data.

# Boundary state x(t) on the facet x1 = x2 (represents permutation (3,2,1))
origin = np.array([2.0, 2.0, 2.0])

# Normal to plane x1 = x2 (facet separating (3,2,1) and (2,3,1))
normal_vec = np.array([1.0, -1.0, 0.0])
normal_vec /= np.linalg.norm(normal_vec)

# Unconstrained -âˆ‡V direction (has both tangential + normal components)
# Think of this as the sorting flow that wants to move toward v_s = (1,2,3).
unconstrained_vec = np.array([1.3, -0.4, 1.0])
unconstrained_vec = unconstrained_vec / np.linalg.norm(unconstrained_vec) * 1.4

# Decompose into tangential + normal
dot_prod         = np.dot(unconstrained_vec, normal_vec)
normal_component = dot_prod * normal_vec
projected_vec    = unconstrained_vec - normal_component

end_unc  = origin + unconstrained_vec   # tip of unconstrained gradient
end_proj = origin + projected_vec       # tip of projected flow
diss_vec = end_unc - end_proj           # the "lost" / dissipated component

# === PLANE PATCH (x1 = x2), SHRUNK AROUND ORIGIN =============================

plane_half = 1  # controls size of plane patch

x_vals = np.linspace(origin[0] - plane_half,
                     origin[0] + plane_half, 4)
z_vals = np.linspace(origin[2] - plane_half,
                     origin[2] + plane_half, 4)

xx, zz = np.meshgrid(x_vals, z_vals)
yy = xx  # still x1 = x2 plane

# === PLOT =====================================================================

fig = plt.figure(figsize=(7, 7))
ax = fig.add_subplot(111, projection='3d')

# Plane style: matching your x1 = x2 permutohedron facet
ax.plot_surface(
    xx, yy, zz,
    alpha=0.15,
    color=SECONDARY_BLUE
)
ax.plot_wireframe(
    xx, yy, zz,
    color=GRAY_DARK,
    linewidth=1.0,
    linestyle='--',
    alpha=0.4
)

# --- Key vertices -------------------------------------------------------------

# 1. Boundary state node (on plane) â€“ represents x_0 with permutation (3,2,1)
ax.scatter(
    origin[0], origin[1], origin[2],
    s=122,
    facecolor=SECONDARY_BLUE,
    edgecolor=BLACK,
    linewidth=1.0,
    zorder=6
)

# 2. Projected tip (on plane) â€“ next infinitesimal step along facet
ax.scatter(
    end_proj[0], end_proj[1], end_proj[2],
    s=122,
    facecolor=SECONDARY_BLUE,
    edgecolor=BLACK,
    linewidth=1.0,
    zorder=6
)

# 3. Unconstrained tip (off plane) â€“ where pure -âˆ‡V wants to go
ax.scatter(
    end_unc[0], end_unc[1], end_unc[2],
    s=122,
    facecolor=GRAY_LIGHT,
    edgecolor=BLACK,
    linewidth=1.0,
    zorder=6
)

# --- Arrows: unconstrained, projected, dissipation ---------------------------

# Unconstrained -âˆ‡V (black, dashed)
arrow_unc = Arrow3D(
    [origin[0], end_unc[0]],
    [origin[1], end_unc[1]],
    [origin[2], end_unc[2]],
    mutation_scale=24,
    lw=6,
    linestyle='--',
    arrowstyle='-|>',
    color="black",
    alpha=0.6,
)
ax.add_artist(arrow_unc)

# Projected flow (black, solid)
arrow_proj = Arrow3D(
    [origin[0], end_proj[0]],
    [origin[1], end_proj[1]],
    [origin[2], end_proj[2]],
    mutation_scale=30,
    lw=6,
    arrowstyle='-|>',
    color="black",
    alpha=1.0,
)
ax.add_artist(arrow_proj)

# Dissipation (normal component) from projected tip to unconstrained tip
arrow_diss = Arrow3D(
    [end_proj[0], end_unc[0]],
    [end_proj[1], end_unc[1]],
    [end_proj[2], end_unc[2]],
    mutation_scale=16,
    lw=6,
    linestyle=':',
    arrowstyle='-|>',
    color="black",
    alpha=0.25,
)
ax.add_artist(arrow_diss)

# ============================================================================
# VERTEX LABELS (positions + text) â€“ symbolic, tied to the paper's story
# ============================================================================

# Positions for vertex labels (kept exactly as you set them)
origin_vert_label_pos = origin   + np.array([-0.1, -0.30, -0.70])
proj_vert_label_pos   = end_proj + np.array([ 0.0,  0.20, -0.10])
unc_vert_label_pos    = end_unc  + np.array([ 0.0,  0.25,  0.10])

# Text for vertex labels
# x_0 corresponds to permutation (3,2,1) and L = [5,4,2] in the example.
origin_vert_text = (
    r'$x_0$' + '\n' +
    r'$(3,2,1)$' + '\n' +
    r'$L = [5,4,2]$'
)

# x_proj is the projected infinitesimal step along the facet
proj_vert_text = r'$x_{\mathrm{proj}}$'

# v_s is the sorted rank word with L_s = [2,4,5]
unc_vert_text = (
    r'$v_s = (1,2,3)$' + '\n' +
    r'$L_s = [2,4,5]$'
)

# Draw vertex labels (same positions, just updated text)
ax.text(
    origin_vert_label_pos[0], origin_vert_label_pos[1], origin_vert_label_pos[2] + .55,
    origin_vert_text,
    color="black",
    fontsize=LABEL_FONT_SIZE,
    ha='center', va='top'
)

ax.text(
    proj_vert_label_pos[0] + .05, proj_vert_label_pos[1] + .05, proj_vert_label_pos[2],
    proj_vert_text,
    color="black",
    fontsize=LABEL_FONT_SIZE,
    ha='center', va='bottom'
)

ax.text(
    unc_vert_label_pos[0], unc_vert_label_pos[1], unc_vert_label_pos[2] + .08,
    unc_vert_text,
    color="black",
    fontsize=LABEL_FONT_SIZE,
    ha='center', va='bottom'
)

# ============================================================================
# EDGE LABELS (positions + text) â€“ separate from vertex labels
# ============================================================================

# Positions for edge labels (kept exactly as you set them)
unc_edge_label_pos  = (origin + end_unc)/2   + np.array([ 0.25,  0.10,  0.10])
proj_edge_label_pos = (origin + end_proj)/2  + np.array([ 0.40, -0.15,  0.00])
diss_edge_label_pos = (end_proj + end_unc)/2 + np.array([-0.70,  0.10,  0.10])

constraint_pos = np.array([
    origin[0] + plane_half + 0.3,
    origin[1],
    origin[2] - plane_half
])

# Text for edge labels â€“ purely mathematical, no invented coordinates
unc_edge_text = r'$-\nabla V$' + '\n' + r'(unconstrained)'

proj_edge_text = r'$\dot{x} = \Pi_{T_{\mathcal{P}_n}(x)}[-\nabla V]$'

diss_edge_text = (
    r'Dissipation' + '\n' +
    r'normal to facet $x_i = x_j$' + '\n' +
    r'$(3,2,1)\leftrightarrow(2,3,1)$'
)

# Draw edge labels
ax.text(
    unc_edge_label_pos[0] - .5, unc_edge_label_pos[1] - .5, unc_edge_label_pos[2],
    unc_edge_text,
    color="black",
    fontsize=LABEL_FONT_SIZE,
    ha='center', va='bottom',
    alpha=0.8
)

ax.text(
    proj_edge_label_pos[0] + .4, proj_edge_label_pos[1] + .2, proj_edge_label_pos[2] - .4,
    proj_edge_text,
    color="black",
    fontsize=LABEL_FONT_SIZE,
    ha='center', va='center',
    alpha=1.0
)

ax.text(
    diss_edge_label_pos[0] + .95, diss_edge_label_pos[1] + .9, diss_edge_label_pos[2],
    diss_edge_text,
    color="black",
    fontsize=LABEL_FONT_SIZE,
    ha='center', va='center',
    alpha=0.6
)

# Constraint boundary label (plane)
ax.text(
    constraint_pos[0], constraint_pos[1], constraint_pos[2],
    r'Constraint Boundary' + '\n' + r'$x_i = x_j$',
    color="black",
    fontsize=LABEL_FONT_SIZE,
    ha='center',
    va='center'
)

# === AXES / 3D GRID ===========================================================

ax.set_xticks([0, 1, 2, 3])
ax.set_yticks([0, 1, 2, 3])
ax.set_zticks([0, 1, 2, 3])

# Hide tick labels for a clean look
ax.set_xticklabels([])
ax.set_yticklabels([])
ax.set_zticklabels([])

# Zoomed-in box so dissipation is clearly visible
ax.set_xlim(1.5, 3.3)
ax.set_ylim(1.5, 3.3)
ax.set_zlim(1.5, 3.3)

ax.grid(True)
ax.xaxis._axinfo["grid"]['color'] = (0.85, 0.85, 0.85, 0.45)
ax.yaxis._axinfo["grid"]['color'] = (0.85, 0.85, 0.85, 0.45)
ax.zaxis._axinfo["grid"]['color'] = (0.85, 0.85, 0.85, 0.45)

ax.xaxis.set_pane_color((1, 1, 1, 0.0))
ax.yaxis.set_pane_color((1, 1, 1, 0.0))
ax.zaxis.set_pane_color((1, 1, 1, 0.0))

ax.view_init(elev=18, azim=135)

plt.tight_layout()

# === SAVE =====================================================================

filename_base = 'tangent_cone_projection_vertex_edge_labels_separate'
plt.savefig(f'{filename_base}.png', dpi=300, bbox_inches='tight')
plt.savefig(f'{filename_base}.pdf', dpi=300, bbox_inches='tight')
plt.savefig(f'{filename_base}.eps', dpi=300, bbox_inches='tight')

input_pdf  = f"{filename_base}.pdf"
output_pdf = f"{filename_base}_scaled.pdf"

target_width_inch  = 5.0
target_height_inch = 5.0
target_w_pt = target_width_inch  * 72
target_h_pt = target_height_inch * 72

try:
    reader = PdfReader(input_pdf)
    writer = PdfWriter()
    page   = reader.pages[0]

    orig_w_pt = float(page.mediabox.width)
    orig_h_pt = float(page.mediabox.height)

    scale_w = target_w_pt / orig_w_pt
    scale_h = target_w_pt / orig_h_pt
    scale   = min(scale_w, scale_h)

    page.scale_by(scale)
    writer.add_page(page)
    with open(output_pdf, "wb") as f:
        writer.write(f)
    print(f"Saved '{output_pdf}' at ~{target_width_inch}\"Ã—{target_height_inch}\".")
except Exception as e:
    print("Skipped PDF scaling:", e)


# %%
