from __future__ import annotations import numpy as np def tetrahedron_design() -> np.ndarray: """Four unit vertices of a regular tetrahedron, a spherical 2-design.""" X = np.array([[1, 1, 1], [1, -1, -1], [-1, 1, -1], [-1, -1, 1]], dtype=float).T return X / np.linalg.norm(X, axis=0, keepdims=True) def slack_matrix(X: np.ndarray, h: np.ndarray, c: float, U: np.ndarray, clamp: bool = False) -> np.ndarray: A = c * np.outer(h, h) - X.T @ U.T @ X return np.maximum(A, 0.0) if clamp else A def polar_attention_logits(qk_logits: np.ndarray, A: np.ndarray, alpha: float = 1.0, eps: float = 1e-8) -> np.ndarray: return qk_logits + alpha * np.log(np.maximum(A, 0.0) + eps) def row_softmax(z: np.ndarray) -> np.ndarray: z = z - np.max(z, axis=-1, keepdims=True) p = np.exp(z) return p / np.sum(p, axis=-1, keepdims=True)