Skip to content

onnx/special: add MeanVarianceNormalization, Hardmax, and hyperbolic ops - #263

Open
axiom-of-choice wants to merge 2 commits into
sbryngelson:mainfrom
axiom-of-choice:feat/onnx-mvn-hardmax-hyperbolic
Open

onnx/special: add MeanVarianceNormalization, Hardmax, and hyperbolic ops#263
axiom-of-choice wants to merge 2 commits into
sbryngelson:mainfrom
axiom-of-choice:feat/onnx-mvn-hardmax-hyperbolic

Conversation

@axiom-of-choice

Copy link
Copy Markdown
Contributor

Three ONNX ops composed from existing graph ops, no new kernels:

  • MeanVarianceNormalization (onnx: MeanVarianceNormalization #123): (x - mean) / sqrt(var + 1e-9) over axes (default [0,2,3]), var = mean(x*x) - mean(x)^2, fixed 1e-9 epsilon.
  • Hardmax (onnx: Hardmax #122): one-hot of the argmax (lowest index on ties) via argmax + select against a baked arange; 2D [C, W] only with axis in {-2,-1,0,1}, raises on higher rank (matches the ANE ArgMax constraint).
  • Sinh/Cosh/Asinh/Acosh/Atanh (onnx: Sinh, Cosh, and the inverse hyperbolics #121): added to aneforge.special (decompositions in exp/log/sqrt) and routed from ONNX; fp16 overflow near |x| ~ 11 documented, tests use the finite range.

Testing

  • Build-level and numeric tests in tests/test_onnx.py (oracle onnxruntime) plus special-function tests in tests/test_special_trig.py (oracle numpy).
  • On-device M2 Pro: tests/test_special_trig.py tests/test_onnx.py → 200 passed. Errors: MVN 2.2e-3 abs; Hardmax exact (0.0); hyperbolics ≤3.4e-3 rel.
  • Off-device suite: 407 passed, 2 skipped.
  • ruff, pylint 2-space gate, pyright aneforge, compileall clean.

@sbryngelson sbryngelson left a comment

Copy link
Copy Markdown
Owner

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Hardmax is solid -- checking the tie cases against onnxruntime was the right call, and it passes on
device. The other two parts need changes.

MeanVarianceNormalization computes variance one-pass as E[x^2] - E[x]^2, which cancels in fp16 as
soon as the mean is large relative to the spread. Measured against the two-pass form on the same
graph, same input, [1,3,4,4] standard normal shifted by a constant:

mean shift    one-pass (this PR)    two-pass
       0.0                0.0018      0.0010
       5.0                0.0463      0.0052
      50.0                   inf      0.0234
     200.0                   inf      0.2921

The test only uses standard_normal, which is the one case where the two agree. Real MVN inputs are
not centered -- that's what the op is for.

The inf is a second, independent bug: var.adds(1e-9) is not a guard, because the ANE flushes
subnormals. Everything below fp16's smallest normal comes back as exactly zero:

eps      fp16 value    sqrt(0 + eps) on ANE
1e-09     0.000e+00                     0.0
1e-07     1.192e-07                     0.0      <- representable as a subnormal, still flushed
1e-06     1.013e-06                     0.0
6e-05     6.002e-05                7.8125e-03
1e-04     1.000e-04                1.0338e-02

So the epsilon has to be at least ~6e-5 to exist at all. Two-pass makes var >= 0 by construction, so
it only matters for genuinely constant input, but it should still be a real number.

On the hyperbolic ops: they pass, but the test metric can't see what it's measuring. abs(out - ref).max() / abs(ref).max()
over linspace(-10, 10) divides by sinh(10) = 11013, so the 5e-3 tolerance permits 55 in absolute
error. Per-point relative error:

sinh(0.001)    2.4% off
atanh(0.003)   4.1% off
asinh(-10)     2.1% off

Only asinh has a clean fix, and I've applied it: asinh is odd, so evaluating on |x| and signing
back avoids the cancellation in log(x + sqrt(x^2+1)) for x << 0. That takes x=-10 from 2.1% to 0.0%
with no regression on positive x.

For sinh and atanh I did not change the formula. expm1 and log1p fix the origin exactly but
are only valid on |x| <= ~0.7 and argument in [-0.5, 1], so using them would need a domain split like
exp_wide/log_wide already do. Instead I documented the actual limit, since the current docstrings
imply overflow near 11 is the only constraint. Happy to take the domain split instead if you'd
rather have the accuracy -- say the word.

Note the patch also updates test_hyperbolic_ops_build, since asinh's terminal op becomes mul.

Verified on device: 203 passed across test_special_trig.py and test_onnx.py, ruff clean. Reverting
just the source changes makes the new tests fail (MVN non-finite at shift 50, asinh 2.08% at x=-10),
so they do pin the behaviour.

diff --git a/aneforge/onnx.py b/aneforge/onnx.py
index be89496..e855994 100644
--- a/aneforge/onnx.py
+++ b/aneforge/onnx.py
@@ -895,8 +895,9 @@ def _mvn(node, ins, a, i):
     return (x - mean) / np.sqrt(var + 1e-9)
   axes = tuple(int(v) % len(x.shape) for v in a.get("axes", [0, 2, 3]))
   mean = x.mean(axes)
-  var = (x * x).mean(axes) - mean * mean
-  return (x - mean) / var.adds(1e-9).sqrt()
+  d = x - mean
+  var = (d * d).mean(axes)                   # two-pass: E[x^2]-E[x]^2 cancels to <=0 in fp16 once mean >> std
+  return d / var.adds(1e-4).sqrt()           # eps below fp16's min normal (6.1e-5) flushes to zero on the ANE
 @onnx_op("Hardmax")
 def _hardmax(node, ins, a, i):
   """Hardmax: one-hot of the argmax (lowest index on ties); 2D [C,W] only, axis in {-2,-1,0,1}."""
diff --git a/aneforge/special.py b/aneforge/special.py
index 375fdf2..4713065 100644
--- a/aneforge/special.py
+++ b/aneforge/special.py
@@ -221,7 +221,11 @@ def log_wide(x: Tensor, sqrts: int = 3) -> Tensor:
 # hyperbolic trig
 
 def sinh(x: Tensor) -> Tensor:
-  """sinh(x) for |x| <= ~10; overflows fp16 near ln(65504) ~ 11 (same wall as softplus)."""
+  """sinh(x) for |x| <= ~10; overflows fp16 near ln(65504) ~ 11 (same wall as softplus).
+
+  Accurate in absolute terms, not relative: e^x - e^-x cancels near the origin, so the relative
+  error reaches ~2.4% for |x| < 0.01 (absolute error stays ~2e-5). Use expm1 if you need the
+  small-argument regime; its Taylor form is only valid for |x| <= ~0.7."""
   return (x.exp() - (x * -1.0).exp()) * 0.5
 
 
@@ -231,8 +235,11 @@ def cosh(x: Tensor) -> Tensor:
 
 
 def asinh(x: Tensor) -> Tensor:
-  """asinh(x) for |x| <= ~10; domain is all real x, but fp16 overflows for large |x|."""
-  return (x + (x * x).adds(1.0).sqrt()).log()
+  """asinh(x) for |x| <= ~10; domain is all real x, but fp16 overflows for large |x|.
+
+  Evaluated on |x| and signed back: asinh is odd, and log(x + sqrt(x^2+1)) cancels for x << 0
+  (2.1% relative error at x=-10 in fp16, vs 0.0% for this form)."""
+  return (x.abs() + (x * x).adds(1.0).sqrt()).log() * x.sign()
 
 
 def acosh(x: Tensor) -> Tensor:
@@ -241,7 +248,11 @@ def acosh(x: Tensor) -> Tensor:
 
 
 def atanh(x: Tensor) -> Tensor:
-  """atanh(x) for |x| < 1; singular at |x| == 1."""
+  """atanh(x) for |x| < 1; singular at |x| == 1.
+
+  Accurate in absolute terms, not relative: (1+x)/(1-x) rounds to ~1 near the origin, so the
+  relative error reaches ~4% for |x| < 0.01 (absolute error stays ~1e-4). log1p fixes the
+  small-argument regime but is only valid for its argument in [-0.5, 1], i.e. |x| <= 1/3."""
   return ((x.adds(1.0)) / (x * -1.0).adds(1.0)).log() * 0.5
 
 
diff --git a/tests/test_onnx.py b/tests/test_onnx.py
index 83d6dc8..a88d485 100644
--- a/tests/test_onnx.py
+++ b/tests/test_onnx.py
@@ -1470,7 +1470,7 @@ def test_negative_axis_matches_positive_equivalent():
 # -- hyperbolic ops, MeanVarianceNormalization, Hardmax --------------------- #
 
 def test_hyperbolic_ops_build():
-  for op, want in [("Sinh", "muls"), ("Cosh", "muls"), ("Asinh", "log"), ("Acosh", "log"), ("Atanh", "muls")]:
+  for op, want in [("Sinh", "muls"), ("Cosh", "muls"), ("Asinh", "mul"), ("Acosh", "log"), ("Atanh", "muls")]:
     m = _model([helper.make_node(op, ["x"], ["y"])], [_vi("x", [1, 4])], [_vi("y", [1, 4])])
     _, out = af.onnx_to_tensor(m); assert out.shape == (1, 4) and out.op == want, f"{op} -> {out.op}"
 
@@ -1487,12 +1487,15 @@ def test_hardmax_build():
   with pytest.raises(NotImplementedError): af.onnx_to_tensor(m)
 
 @requires_ane
-def test_mvn_numeric():
+@pytest.mark.parametrize("shift", [0.0, 5.0, 50.0])
+def test_mvn_numeric(shift):
+  """Shifted means are the fp16 trap: one-pass E[x^2]-E[x]^2 cancels and divides by zero at shift=50."""
   pytest.importorskip("onnxruntime")
-  rng = np.random.default_rng(40); x = rng.standard_normal((1, 3, 4, 4)).astype(np.float32)
+  rng = np.random.default_rng(40); x = (rng.standard_normal((1, 3, 4, 4)) + shift).astype(np.float32)
   m = _model([helper.make_node("MeanVarianceNormalization", ["x"], ["y"])], [_vi("x", [1, 3, 4, 4])], [_vi("y", [1, 3, 4, 4])])
   got, ref = _run_vs_ort(m, x)
-  assert np.abs(got - ref).max() < 1e-2
+  assert np.isfinite(got).all(), f"MVN produced non-finite output at mean shift {shift}"
+  assert np.abs(got - ref).max() < 5e-2
 
 @requires_ane
 def test_hardmax_numeric():
diff --git a/tests/test_special_trig.py b/tests/test_special_trig.py
index 9b138bf..2c7b293 100644
--- a/tests/test_special_trig.py
+++ b/tests/test_special_trig.py
@@ -106,6 +106,17 @@ def test_asinh_decomposition_matches_numpy():
   assert np.abs(out - ref).max() < 1e-1
 
 
+def test_asinh_is_odd_and_accurate_for_negative_x():
+  """Per-point relative error, which an error-over-max-|ref| metric cannot see: the naive
+  log(x + sqrt(x^2+1)) cancels for x << 0 and is 2.1% off at x=-10."""
+  x = np.array([[-10.0, -6.0, -3.0, -1.0, 1.0, 3.0, 6.0, 10.0]], np.float16)
+  out = _run(special.asinh, x)
+  ref = np.arcsinh(x.astype(np.float32))
+  rel = np.abs(out - ref) / np.abs(ref)
+  assert rel.max() < 5e-3, f"asinh per-point relerr {rel.max():.2%} at x={float(x[0][rel.argmax()])}"
+  assert np.abs(out + out[:, ::-1]).max() < 1e-3, "asinh must be odd"
+
+
 def test_acosh_decomposition_matches_numpy():
   x = np.linspace(1.0001, 10.0, 128).astype(np.float16).reshape(1, 128)
   out = _run(special.acosh, x)

@axiom-of-choice

Copy link
Copy Markdown
Contributor Author

Applied the review patch, plus one adjustment that the on-device run forced.

What's in the branch now (cb7d153):

  • MVN: two-pass variance (d = x - mean; var = (d*d).mean) and 1e-4 epsilon (below fp16 min normal the ANE flushes to zero).
  • asinh: evaluated on |x| and signed back; sinh/atanh get the accuracy warnings in their docstrings, formulas unchanged.
  • Tests updated: test_mvn_numeric parametrized over shifts, test_hyperbolic_ops_build Asinh expectation log -> mul, plus test_asinh_is_odd_and_accurate_for_negative_x.

Adjustment: test_mvn_numeric[50.0] fails on my M2 Pro at the 5e-2 tolerance — max abs error 0.1052 (seed 40). This is the fp16 quantization floor of the shifted input, not a regression: ulp(50) in fp16 is ~0.049, so the centered x - mean carries ~5% input noise that lands on the output. Measured max over 16 seeds at shift 50: 0.1052 (median 0.066, p95 0.0995). Your M5 passes at 5e-2 because the ANE fp16 rounding differs; M2 Pro is the worst case here.

I loosened the tolerance to 1.5e-1 (above the measured max) and kept the isfinite assertion, which is the assertion that actually pins the one-pass inf bug — reverting the source change still makes it fail (non-finite at shift 50).

Verified on device (M2 Pro): tests/test_special_trig.py tests/test_onnx.py -> 203 passed. Off-device: 407 passed, 2 skipped. ruff, pylint 10/10 clean.

Note: the PR body still says "MVN 2.2e-3 abs", which no longer reflects the shifted-input behavior; error is bounded by ~1e-1 at shift 50, ~3e-3 near the origin.

@sbryngelson sbryngelson left a comment

Copy link
Copy Markdown
Owner

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Taking the tolerance change -- you're right and I was wrong to pin it at 5e-2. That number was
fitted to the machine I measured on and had no headroom. On M5:

shift= 0.0   max abs err 0.0012    fp16 ulp 0.0010
shift= 5.0   max abs err 0.0081    fp16 ulp 0.0039
shift=50.0   max abs err 0.0486    fp16 ulp 0.0312

0.0486 against a 0.05 bound would have gone red on someone else's hardware sooner or later, so
loosening it was the right call.

Two small follow-ups, neither blocking, and I'm happy to merge without them.

The loose bound costs some regression signal. I reverted the MVN fix and re-ran your updated test:
it still fails, but only on the isfinite assert at shift 50. The shift-5 case, which the one-pass
formula gets wrong at 0.0659, now passes under 1.5e-1. The ulp argument only justifies the loose
bound at shift 50, so a per-shift tolerance would keep the other two tight:

@pytest.mark.parametrize("shift,tol", [(0.0, 5e-3), (5.0, 2e-2), (50.0, 1.5e-1)])
def test_mvn_numeric(shift, tol):

And the comment says "fp16 ulp at shift 50 is ~0.049" -- 0.049 is the measured error; the ulp there
is 0.0312 (np.spacing(np.float16(50))). The conclusion is unchanged, but this repo leans on its
numeric comments being exact.

Verified on device: 203 passed across test_special_trig.py and test_onnx.py, ruff clean.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants