"""loss_forms.py — the composable loss library of the loss campaign. #TAG:loss_forms #TAG:accumulation Every RUNNABLE loss form in one place: the four differencing primitives, the accumulation formats A0-A8 as composable functions, the candidate losses (FAC, PWA weights, compartment roles, latent-chain), and a self-smoke. Deliberately ABSENT, by statute (inventory/LOSS_MANIFEST.md): A9 sum-no-norm — scale rides on batch/seq; lr stops transferring. A10 EMA/cross-step — the VQ/commitment/load-balancing failure class. InfoNCE into address paths — legal only as a readout head (L-113 / L-017). House laws honored throughout: pure Adam wd=0 (constructor not included here — use amoe.laws.make_optimizer); fp32/TF32-off; crc32 seeds never hash(); CV is a readout never a force; masking never renormalizes; gauges fp64. Colab-cell-safe: no argparse side effects, no __file__ logic. Smoke: python tools/loss_forms.py """ import math import os import sys import zlib import torch import torch.nn.functional as F def _root(): d = os.path.abspath(os.getcwd()) while True: if os.path.exists(os.path.join(d, "MANIFEST.md")): return d p = os.path.dirname(d) if p == d: return os.getcwd() d = p for _p in (os.path.join(_root(), "tools"),): if os.path.isdir(_p) and _p not in sys.path: sys.path.insert(0, _p) def seed_for(name: str) -> int: return zlib.crc32(name.encode("utf-8")) & 0x7FFFFFFF # ============================================================= PRIMITIVES # Each returns PER-ELEMENT residuals (unreduced) so accumulation composes. def prim_ce(logits, target): """CE — the coupled primitive: log-sum-exp partition over the last dim. Hessian diag(p)-pp^T: exact null direction; spectrum collapses as p_max->1 (measured T21). Returns (...,) per-position nats.""" return F.cross_entropy(logits.reshape(-1, logits.shape[-1]), target.reshape(-1), reduction="none" ).reshape(target.shape) def prim_sq(pred, target): """Squared error — the aleph's only sanctioned codebook pressure rides this (recon through M-hat). Returns per-element squares.""" return (pred - target) ** 2 def prim_kl(logits, teacher_probs): """KL to detached teacher probs. LAW: alpha <= 0.25, NEVER on founders, never in a selection loop without a quality gap (L-114/L-022).""" return F.kl_div(F.log_softmax(logits, -1), teacher_probs.detach(), reduction="none").sum(-1) def prim_cosh_bregman(v, code, mu=1.0, clamp=4.0): """BREG — the Bregman divergence of the aleph's own potential sum-cosh: D_Phi(v - c*mu, 0) = cosh(r) - 1. Uncoupled per axis, curvature >= 1, antipodally invariant L(v,c)==L(-v,-c). VERDICT ON RECORD (L-070/L-138): loses to CE wherever CE is healthy, 3/3; DECOMPRESSES the coupled- partition collapse 3/3 (usage 1-2.7 -> ~61/64). Use it where the partition coupling is the disease, not as a general replacement.""" r = (v - code * mu).clamp(-clamp, clamp) return torch.cosh(r) - 1.0 # ====================================================== ACCUMULATION FORMATS # Each takes per-element residuals -> a scalar (or a weighted scalar). def a0_mean(res): """A0 uniform-mean. The default; honestly dominant (71/138).""" return res.mean() def a1_chunked_ce(logits_fn, hidden, target, chunk=512, ignore_index=-100): """A1 chunk-sum-renormalize for CE: never materialize seq x vocab. Mathematically identical to A0; a 5x memory law (L-004). `logits_fn` maps a hidden slice -> logits (the lm head).""" s, n = None, 0 T = hidden.shape[1] for i in range(0, T, chunk): lg = logits_fn(hidden[:, i:i + chunk]) tgt = target[:, i:i + chunk] term = F.cross_entropy(lg.reshape(-1, lg.shape[-1]), tgt.reshape(-1), ignore_index=ignore_index, reduction="sum") s = term if s is None else s + term n += int((tgt != ignore_index).sum()) return s / max(n, 1) def a2_weighted(res, w, dims=None): """A2 per-sample(or per-element)-then-weighted. Reduce res over `dims` FIRST if given, then weight and renormalize by w.sum() — a mean taken too early silently erases the weight.""" if dims is not None: res = res.mean(dim=dims) return (res * w).sum() / w.sum().clamp_min(1e-12) def a3_band_composed(res_per_band, w_bands): """A3 band-crossfade composition: (B, N_BANDS) losses x (B, N_BANDS) windows -> scalar. Windows must be a partition of unity on the TRAINING coordinate (band coordinate law); isolation is quadratic in the window.""" return (res_per_band * w_bands).sum(-1).mean() def a4_masked(res, mask): """A4 masked-denominator — THE SILENT-ZERO CLASS. Asserts the mask fired: a term that never fires is indistinguishable from a null term.""" m = mask.float() live = m.sum() assert float(live) > 0, "A4 silent zero: mask never fired (assert the count upstream)" return (res * m).sum() / live def a5_dose_coupled(base_scalar, aux_res, w_route, lam=1.0): """A5 dose-coupled auxiliary: base + lam * routed aux. lam~1 is the measured operating point on the flow substrate; run the CONDITIONING GATE on the aux's recovery map before spending (L-016 vs L-115).""" return base_scalar + lam * a2_weighted(aux_res, w_route) def a6_paired(res_a, res_b): """A6 paired-difference: identical (row, noise, t) triples per arm, per-sample reduction, fp64 accumulation. Without pairing, sub-1% effects are invisible (the ~0.988 unpaired floor).""" return (res_a.double() - res_b.double()).mean() def a7_grid_infonce(za, zb, temp=0.07): """A7 grid-pairwise (InfoNCE), symmetric. THE LOUDEST GRADIENT — legal ONLY as a readout objective on a head outside the compute path; NEVER into address paths (L-113). You are responsible for that placement.""" sims = za @ zb.t() / temp lbl = torch.arange(za.shape[0], device=za.device) return (F.cross_entropy(sims, lbl) + F.cross_entropy(sims.t(), lbl)) / 2 def a8_fp64_gauge(fn, *args): """A8 fp64-accumulate for GAUGES (no_grad, autocast off). fp32 CM dets lose ~4% on near-degenerate pentachora.""" with torch.no_grad(): return fn(*(a.double() if torch.is_tensor(a) else a for a in args)) # ========================================================== CANDIDATE LOSSES def fac_loss(feats, R, code_rows, mu=1.0, t_loss=0.3): """FAC: normalize(feats) @ R^T / t_loss -> cosh-Bregman to the target code. R is a FROZEN orthonormal frame (gauge-fixed by construction); code_rows in {-1,+1}^K frozen. See prim_cosh_bregman's verdict note.""" v = (F.normalize(feats, dim=-1) @ R.t()) / t_loss return prim_cosh_bregman(v, code_rows, mu=mu) def pwa_weights(pi, form="inverse", w_min=0.1, band=(0.10, 0.60), eps=0.02): """PWA weight builders over a FROZEN reference's true-token prob pi. GATE RECORD (2026-07-25, trained-ce reference): band-kernel novelty 0.0145 REFUSED; window 0.0562 marginal; inverse 0.0832 weak-pass — all far below the 0.715 payer class. CONDITIONAL: do not spend an arm matrix on these; revival bar is a form with novelty >= 0.3.""" if form == "band-kernel": return w_min + (1 - w_min) * 4 * pi * (1 - pi) if form == "window": lo, hi = band return (torch.sigmoid((pi - lo) / eps) * torch.sigmoid((hi - pi) / eps)).clamp_min(w_min) if form == "inverse": return (1 - pi).clamp_min(w_min) raise ValueError(form) # Compartment ROLE losses (rank 1 of the series). Each supervises a DIFFERENT # QUANTITY through the band's channel window — the 0.715-class design contract # (a reweighting of the base residual would be gate-refused; these are not). # `cmap` is compartment_smoke.build_compartment_map(...); h is the trunk # hidden (B, T, d). Fixed probes are frozen buffers (placement by # construction); trainable role heads replace them in a real bed. def role_low_recon(h, W_chan, emb_target, probe): """LOW = absolute/reconstructive: rebuild the token's own input embedding from the LOW channels alone. The aleph's proven pressure class.""" hw = h * W_chan[:, 0] return prim_sq(hw @ probe, emb_target.detach()).mean(-1) def role_mid_continuity(h, W_chan, probe): """MID = relational: geodesic continuity of adjacent-position MID-channel states (1 - cos on a fixed projection). A different quantity (the trajectory), not a reweighting of the next-token residual.""" z = F.normalize((h * W_chan[:, 1]) @ probe, dim=-1) return 1.0 - (z[:, :-1] * z[:, 1:]).sum(-1) def role_high_span(h, W_chan, span_target, probe, span=32): """HIGH = structural: predict the span's byte-histogram signature from the HIGH channels. Span pooling over TIME toward an explicit span-level TARGET (not GAP-in-an-encoder: the pooled object IS the supervised quantity, flagged per the GAP law regardless).""" B, T, d = h.shape n = T // span hw = (h * W_chan[:, 2])[:, :n * span].reshape(B, n, span, d).mean(2) return prim_sq(hw @ probe, span_target.detach()).mean(-1) def latent_chain_terms(feats_answer, feats_register, R, code_y, code_z, mu=1.0, t_loss=0.3, lam=1.0): """LATENT-CHAIN: FAC on the answer position + FAC on a LATENT register position targeting the intermediate value's code — supervision of a quantity NOT in the output string (the thing CE structurally cannot express). Mandatory control in any bed: latent_chain_shuffled (c_z drawn from a shuffled intermediate). Prereg: direct composite 0.0 -> >= 0.50, REFUTED < 0.10.""" la = fac_loss(feats_answer, R, code_y, mu, t_loss) lz = fac_loss(feats_register, R, code_z, mu, t_loss) return la.mean() + lam * lz.mean() # ================================================= LEGACY ROSTER (extracted) # Every historical form with a recorded formula and no living local impl, # made runnable. Verdicts travel in the docstrings; the manifest row is the # authority (inventory/LOSS_MANIFEST.md). def margin_head(feats, weight, target, kind="arcface", s=30.0, m=0.30): """L-036 RoseFace margin family. cos(th+m) (arc) | cos(th)-m (cos) | cos(m*th) (sphere), scale s. Historical ceiling: 60% single-stream (diagnosed as frozen pentachora + erosion, not the margin).""" z = F.normalize(feats, dim=-1) @ F.normalize(weight, dim=-1).t() th = torch.arccos(z.clamp(-1 + 1e-7, 1 - 1e-7)) if kind == "arcface": zt = torch.cos(th + m) elif kind == "cosface": zt = z - m elif kind == "sphereface": zt = torch.cos(m * th) else: raise ValueError(kind) logits = z.clone() logits.scatter_(-1, target.unsqueeze(-1), zt.gather(-1, target.unsqueeze(-1))) return prim_ce(s * logits, target) def cv_band_loss(anchors, cv_target=0.20, weight=1e-3, n_sets=64, seed=0): """L-040 — THE ONE SANCTIONED CV FORCE. Arm-gated by statute: weight HARD CEILING 1e-3; S^15-class BANKS only, NEVER the aleph codebook; forward loss; fp64 determinant; fixed-seed subset draw (deterministic across steps). Port of tools/exp017_aleph_constellation.py:154-186.""" assert weight <= 1e-3, "CV force above 1e-3 is prohibited (L-110)" A = F.normalize(anchors, dim=-1) n = A.shape[0] assert n >= 5, "pentachoron CV needs >= 5 anchors" g = torch.Generator(device="cpu").manual_seed(seed) idx = torch.stack([torch.randperm(n, generator=g)[:5] for _ in range(n_sets)]) pts = A[idx] d2 = torch.cdist(pts.double(), pts.double()).pow(2) cm = torch.ones(n_sets, 6, 6, dtype=torch.float64, device=A.device) cm[:, 0, 0] = 0.0 cm[:, 1:, 1:] = d2 v = (-torch.linalg.det(cm) / 9216.0).clamp_min(1e-24).sqrt() cv = (v.std() / v.mean().clamp_min(1e-12)).float() return weight * (cv - cv_target).abs() def cm_validity_hinge(pts, lam=0.01, eps=1e-6): """L-045 KSimplex validity hinge: penalize non-positive CM volume^2 on the simplex. Requires d/k >= 8 or the det is numerically unstable.""" B = pts.shape[0] d2 = torch.cdist(pts, pts).pow(2) k1 = pts.shape[1] cm = torch.ones(B, k1 + 1, k1 + 1, dtype=pts.dtype, device=pts.device) cm[:, 0, 0] = 0.0 cm[:, 1:, 1:] = d2 sign = -1.0 if (k1 % 2 == 0) else 1.0 vol2 = sign * torch.linalg.det(cm) return lam * F.relu(eps - vol2).mean() def cm_volume_spread(vol2_per_layer, lam=0.005): """L-046 volume-spread REWARD: -std(log|vol^2|) across layers — an anti-collapse diversity reward, note the SIGN.""" return -lam * torch.log(vol2_per_layer.abs().clamp_min(1e-24)).std() def procrustes_sq(A, B): """L-047/L-111 Procrustes residual ||A R* - B||^2 (R* via SVD). PLACEMENT VERDICT: as a x0.3 regularizer beside a real force it tightens CV (rating 6); as THE training force R@1 = 0.000 (rating 1). It measures alignability; it cannot create it.""" U, _, Vt = torch.linalg.svd(A.t() @ B) R = U @ Vt return ((A @ R - B) ** 2).mean() def soft_hand_weights(cv, target, sigma=0.15, boost=1.5, penalty=1.0): """L-026 soft hand — reward, not penalty: near the CV target the recon gradient is BOOSTED (1..1+boost); far, a restoring force. Adverse finding on record: SUSTAINED moderate boost hurts (the model optimizes for staying in the boost zone). Returns (recon_weight, cv_penalty).""" prox = torch.exp(-((cv - target) ** 2) / (2 * sigma ** 2)) return 1.0 + boost * prox, penalty * (1.0 - prox) def kd_guard(alpha, is_founder=False, in_selection_loop=False, teacher_gap=None): """L-022/L-114 KD statute: alpha <= 0.25, never on founders, never in a selection loop without a quality gap. Raises on the L-114 configuration (inverse evolution, 2.4301 -> 2.5603).""" if is_founder: raise ValueError("KD on a founder is prohibited (L-114)") if alpha > 0.25 and in_selection_loop and not teacher_gap: raise ValueError("KD alpha > 0.25 in a selection loop without a " "quality gap reproduces inverse evolution (L-114)") return min(alpha, 1.0) # ============================================= DEVIANT ROSTER (gate-cleared) # inventory/DEVIANT_ROSTER.md candidates. Novelty numbers travel with them; # trained verdicts graduate them to LOSS_MANIFEST rows. def dev_softmax_accum(res, T=0.5): """Worst-position accumulation: T*logsumexp(res/T) - T*log(N). Gradient == softmax(res/T) weighting (self-paced weighting IS this loss). Gate 0.911 at trained state - the highest ever. FLAG: on natural text the worst positions are largely irreducible entropy; prereg carries a held-out bar.""" flat = res.reshape(-1) return T * torch.logsumexp(flat / T, 0) - T * math.log(flat.numel()) def dev_geomean_accum(res, eps=1e-3): """Geometric-mean accumulation: mean(log(res+eps)) - the anti-focal (gradient 1/res polishes the nearly-solved). Gate 0.486 trained.""" return torch.log(res + eps).mean() def sparsemax_loss(z, y): """Sparsemax loss (Martins & Astudillo 2016): a PARTIAL partition - sparse support - between CE (full coupling) and FAC (zero coupling). The coupling-axis probe for the L-138 mechanism. Gate 0.253 (state- independent form). z: (N,V) logits, y: (N,) targets -> (N,) losses.""" zs, _ = torch.sort(z.detach(), dim=-1, descending=True) cs = zs.cumsum(-1) k = torch.arange(1, z.shape[-1] + 1, device=z.device, dtype=z.dtype) ksup = ((1 + k * zs) > cs).to(z.dtype).sum(-1, keepdim=True) tau = (cs.gather(-1, ksup.long() - 1) - 1) / ksup psp = (z - tau).clamp_min(0) # sparsemax probs (grad ok) zy = z.gather(-1, y.unsqueeze(-1)).squeeze(-1) zsq = torch.where(psp > 0, z ** 2 - tau ** 2, torch.zeros_like(z)).sum(-1) return -zy + 0.5 * zsq + 0.5 def fac_loss_link(feats, R, code_rows, link="cosh", mu=1.0, t_loss=0.3): """The FAC link dial: cosh (exponential tails, the measured verdict) | tanh-Hamming (bounded) | cauchy log(1+r^2) (sub-quadratic). Links are ~90% collinear at init (direction dominates early; tails matter late).""" v = (F.normalize(feats, dim=-1) @ R.t()) / t_loss if link == "cosh": return prim_cosh_bregman(v, code_rows, mu=mu) if link == "tanh": return 1.0 - torch.tanh(v) * code_rows if link == "cauchy": return torch.log1p((v - code_rows * mu) ** 2) raise ValueError(link) # ================================================ FORBIDDEN CONTROLS [FORCE] # Runnable ONLY as explicitly-forced control arms (the blob-on-eps pattern: # the library refuses the design and permits the falsification). Each cites # its manifest row and warns loudly. def _force_gate(force, row, evidence): if not force: raise ValueError( f"{row} is a FORBIDDEN class ({evidence}). This implementation " f"exists ONLY as a control arm - pass force=True to reproduce " f"the failure on purpose.") import warnings warnings.warn(f"{row} forced: you are reproducing a documented failure " f"class as a CONTROL, not training a design.") def forbidden_vq_commitment(z_e, codebook, beta=0.25, force=False): """L-105 VQ codebook + commitment loss (EMA variant NOT provided — the cross-step state is A10 and stays absent even here). Evidence: the aleph codebook holds 125+/128 axes alive at div_weight=0 without it.""" _force_gate(force, "L-105 VQ/commitment", "14x path collapse class") d = torch.cdist(z_e.reshape(-1, z_e.shape[-1]), codebook) e = codebook[d.argmin(-1)].reshape(z_e.shape) return (prim_sq(z_e.detach(), e).mean() + beta * prim_sq(z_e, e.detach()).mean()) def forbidden_load_balancing(router_probs, expert_mask, alpha=0.01, force=False): """L-134 switch-style balance aux: alpha * N * sum_i f_i * P_i. Evidence: banned and never needed — usage stays near-uniform read-only.""" _force_gate(force, "L-134 load-balancing aux", "no-balancing statute") N = router_probs.shape[-1] f = expert_mask.float().mean(dim=tuple(range(expert_mask.ndim - 1))) P = router_probs.mean(dim=tuple(range(router_probs.ndim - 1))) return alpha * N * (f * P).sum() def forbidden_gap(x, spatial_dims, force=False): """L-109 global average pooling in a geometric encoder. Evidence: 70% -> 29% collapse, replicated twice. Patch aggregation defaults to MEAN over tokens at the READOUT, never pooling inside the encoder.""" _force_gate(force, "L-109 GAP", "70->29 collapse, replicated") return x.mean(dim=spatial_dims) # ================================================================ THE GATES def collinearity_novelty(loss_arm, loss_base, params): """novelty = 1 - |cos(grad_arm, grad_base)|. Composed role arms are judged whole; additive auxiliaries are judged as THE TERM BEING ADDED. Calibration: HP/LP 0.0026-0.0083 (inert) vs blob 0.715 (payer). REFUSE below 0.05; the payer class starts ~0.3.""" ga = torch.autograd.grad(loss_arm, params, retain_graph=True, allow_unused=True) gb = torch.autograd.grad(loss_base, params, retain_graph=True, allow_unused=True) # zero-fill on the SHARED parameter support: a param an arm does not # touch contributes the zero vector to its direction (dropping it would # misalign the two flattened gradients) fa = torch.cat([(g if g is not None else torch.zeros_like(p)).reshape(-1) for g, p in zip(ga, params)]) fb = torch.cat([(g if g is not None else torch.zeros_like(p)).reshape(-1) for g, p in zip(gb, params)]) return 1.0 - abs(F.cosine_similarity(fa.unsqueeze(0), fb.unsqueeze(0)).item()) # ================================================================ SELF-SMOKE def _smoke(): torch.backends.cuda.matmul.allow_tf32 = False torch.backends.cudnn.allow_tf32 = False dev = "cuda" if torch.cuda.is_available() else "cpu" if dev == "cuda": torch.cuda.set_per_process_memory_fraction(0.73) g = torch.Generator().manual_seed(seed_for("loss_forms")) ok = [] B, T, V, d, K = 4, 64, 256, 192, 64 logits = torch.randn(B, T, V, generator=g, requires_grad=True) y = torch.randint(0, V, (B, T), generator=g) ce = prim_ce(logits, y) ok.append(("prim_ce shape+grad", bool(ce.shape == (B, T) and torch.autograd.grad(ce.mean(), logits)[0].abs().sum() > 0))) # A1 == A0 identity (the 5x memory law is math-free) h = torch.randn(B, T, d, generator=g) W = torch.randn(V, d, generator=g) * 0.02 fn = lambda hh: hh @ W.t() full = F.cross_entropy(fn(h).reshape(-1, V), y.reshape(-1)) ok.append(("A1 == A0 exactly", torch.allclose(a1_chunked_ce(fn, h, y, chunk=17), full, atol=1e-6))) # A2 early-mean hazard: weighting after full mean == unweighted res = torch.randn(B, T, generator=g).abs() w = torch.rand(B, generator=g) + 0.1 good = a2_weighted(res, w, dims=(1,)) bad = res.mean() * (w / w).mean() ok.append(("A2 weight not erased", abs(good - res.mean()) > 1e-6 and torch.allclose(bad, res.mean()))) # A4 silent-zero assert fires try: a4_masked(res, torch.zeros_like(res)) ok.append(("A4 silent-zero assert", False)) except AssertionError: ok.append(("A4 silent-zero assert", True)) # A6 fp64; A7 symmetric ok.append(("A6 fp64", a6_paired(res, res).dtype == torch.float64 and float(a6_paired(res, res)) == 0.0)) za = F.normalize(torch.randn(8, 32, generator=g), dim=-1) zb = F.normalize(torch.randn(8, 32, generator=g), dim=-1) ok.append(("A7 symmetric", torch.allclose(a7_grid_infonce(za, zb), a7_grid_infonce(zb, za), atol=1e-6))) # FAC: antipodal invariance + gradient flow through feats feats = torch.randn(B, T, 256, generator=g, requires_grad=True) R = torch.linalg.qr(torch.randn(256, 256, generator=g))[0][:K] code = ((torch.randn(V, K, generator=g) > 0).float() * 2 - 1)[y] L = fac_loss(feats, R, code).mean() v = (F.normalize(feats, dim=-1) @ R.t()) / 0.3 ok.append(("FAC antipodal + grad", bool(torch.allclose(prim_cosh_bregman(v, code), prim_cosh_bregman(-v, -code)) and torch.autograd.grad(L, feats)[0].abs().sum() > 0))) # PWA weights bounded + floored pi = torch.rand(B, T, generator=g) for f in ("band-kernel", "window", "inverse"): wf = pwa_weights(pi, f) ok.append((f"PWA {f} in [w_min,1]", float(wf.min()) >= 0.1 - 1e-6 and float(wf.max()) <= 1.0 + 1e-6)) # Compartment roles: shapes + grad + zero-grad outside their window try: from compartment_smoke import build_compartment_map cmap = build_compartment_map(P=32, Ds=4, d=d) Wc = cmap["W_chan_band"] hh = torch.randn(B, T, d, generator=g, requires_grad=True) pl = torch.randn(d, 48, generator=g) / math.sqrt(d) emb_t = torch.randn(B, T, 48, generator=g) lo = role_low_recon(hh, Wc, emb_t, pl).mean() gl = torch.autograd.grad(lo, hh)[0] dead = (Wc[:, 0] == 0) ok.append(("role LOW grad confined to LOW channels", bool(float(gl[..., dead].abs().sum()) == 0.0 and float(gl.abs().sum()) > 0))) mid = role_mid_continuity(hh, Wc, pl).mean() sp_t = torch.randn(B, T // 32, 48, generator=g) hi = role_high_span(hh, Wc, sp_t, pl).mean() ok.append(("roles MID/HIGH finite+grad", bool(torch.isfinite(mid) and torch.isfinite(hi) and torch.autograd.grad(mid + hi, hh)[0].abs().sum() > 0))) except ImportError: ok.append(("compartment roles (map import)", None)) # legacy roster W2 = torch.randn(10, 64, generator=g) f2 = torch.randn(6, 64, generator=g, requires_grad=True) y2 = torch.randint(0, 10, (6,), generator=g) mh = margin_head(f2, W2, y2, "arcface").mean() ok.append(("margin_head grad + finite", bool(torch.isfinite(mh) and torch.autograd.grad(mh, f2)[0].abs().sum() > 0))) bank = torch.randn(96, 16, generator=g, requires_grad=True) cvl = cv_band_loss(bank) ok.append(("cv_band_loss forward+grad, ceiling enforced", bool(torch.isfinite(cvl) and torch.autograd.grad(cvl, bank)[0].abs().sum() > 0))) try: cv_band_loss(bank.detach(), weight=1e-2) ok.append(("cv_band_loss ceiling assert", False)) except AssertionError: ok.append(("cv_band_loss ceiling assert", True)) pts5 = torch.randn(8, 5, 32, generator=g, requires_grad=True) hinge = cm_validity_hinge(pts5) ok.append(("cm_validity_hinge finite", bool(torch.isfinite(hinge)))) ok.append(("cm_volume_spread sign is a reward", bool(cm_volume_spread(torch.rand(6, generator=g) + 0.1) <= 0))) A2m = torch.randn(32, 8, generator=g); B2m = torch.randn(32, 8, generator=g) ok.append(("procrustes_sq beats unaligned", bool(procrustes_sq(A2m, B2m) <= ((A2m - B2m) ** 2).mean() + 1e-5))) rw, cp = soft_hand_weights(torch.tensor(0.20), 0.20) ok.append(("soft_hand at target: boost on, penalty ~0", bool(rw > 2.4 and cp < 1e-6))) try: kd_guard(0.5, is_founder=True) ok.append(("kd_guard founder refusal", False)) except ValueError: ok.append(("kd_guard founder refusal", True)) # forbidden controls refuse without force, run with it ze = torch.randn(4, 7, 16, generator=g); cb = torch.randn(32, 16, generator=g) import warnings refuse = 0 for fn, args in ((forbidden_vq_commitment, (ze, cb)), (forbidden_load_balancing, (torch.softmax(torch.randn(64, 8, generator=g), -1), F.one_hot(torch.randint(0, 8, (64,), generator=g), 8))), (forbidden_gap, (torch.randn(2, 3, 8, 8, generator=g), (2, 3)))): try: fn(*args) except ValueError: refuse += 1 with warnings.catch_warnings(): warnings.simplefilter("ignore") out = fn(*args, force=True) refuse += int(bool(torch.isfinite(out if out.dim() == 0 else out.sum()))) ok.append(("forbidden controls: refuse w/o force, run with it", refuse == 6)) npass = sum(1 for _, v in ok if v is True) nfail = sum(1 for _, v in ok if v is False) print("LOSS_FORMS SELF-SMOKE") for name, v in ok: print(" %-38s %s" % (name, "PASS" if v is True else ("SKIP" if v is None else "FAIL"))) print("PASS %d FAIL %d SKIP %d" % (npass, nfail, len(ok) - npass - nfail)) return nfail == 0 if __name__ == "__main__": sys.exit(0 if _smoke() else 1)