From f196af2327ceed2eeaef0b4d4248d479180a0087 Mon Sep 17 00:00:00 2001 From: chenyu Date: Fri, 28 Aug 2026 13:47:47 -0400 Subject: [PATCH] fix onnx MeanVarianceNormalization arg (#17812) axes, not axis --- test/external/external_test_onnx_ops.py | 6 ++++++ tinygrad/nn/onnx.py | 9 ++++----- 2 files changed, 10 insertions(+), 5 deletions(-) diff --git a/test/external/external_test_onnx_ops.py b/test/external/external_test_onnx_ops.py index 55256f153a..02071da695 100644 --- a/test/external/external_test_onnx_ops.py +++ b/test/external/external_test_onnx_ops.py @@ -54,6 +54,12 @@ class TestMainOnnxOps(TestOnnxOps): outputs = ["squeezed"] self.helper_test_single_op("Squeeze", inputs, attributes, outputs) + def test_mean_variance_normalization_axes(self): + inputs = {"x": np.random.randn(2, 3, 4, 5).astype(np.float32)} + attributes = {"axes": [2, 3]} + outputs = ["out"] + self.helper_test_single_op("MeanVarianceNormalization", inputs, attributes, outputs) + def test_conv(self): # test VALID auto_pad inputs = { diff --git a/tinygrad/nn/onnx.py b/tinygrad/nn/onnx.py index 7723e74b3e..bc8ca04f82 100644 --- a/tinygrad/nn/onnx.py +++ b/tinygrad/nn/onnx.py @@ -937,9 +937,8 @@ def get_onnx_ops() -> dict[str, types.FunctionType|dict[OpSetId, types.FunctionT if segment_embedding is not None: embedding_sum = embedding_sum + embedding(segment_ids, segment_embedding.shape[0], segment_embedding) out = embedding_sum.layernorm(eps=epsilon) * gamma + beta return out, None, embedding_sum - def MeanVarianceNormalization(x:Tensor, axis:Sequence[int]|None=None): - if axis is None: axis = [0,2,3] - return (x - x.mean(axis, keepdim=True)) / (x.std(axis, keepdim=True, correction=0) + 1e-9) + def MeanVarianceNormalization(x:Tensor, axes:Sequence[int]=(0,2,3)): + return (x - x.mean(axes, keepdim=True)) / (x.std(axes, keepdim=True, correction=0) + 1e-9) def LpNormalization(x:Tensor, axis:int=-1, p:int=2): return x / (x.abs().sum(axis, keepdim=True) if p == 1 else x.square().sum(axis, keepdim=True).sqrt()) @@ -1246,8 +1245,8 @@ def get_onnx_ops() -> dict[str, types.FunctionType|dict[OpSetId, types.FunctionT ret = _qlinearop_float(GlobalAveragePool, [X], [x_zero_point], [x_scale], y_scale, y_zero_point) return ret.permute(0, *range(2, ret.ndim), 1) if channels_last else ret # NCHW -> NHWC - def ConvInteger(x: Tensor, w: Tensor, x_zero_point:Tensor = Tensor(0), w_zero_point:Tensor = Tensor(0), B: Tensor | None = None, **opts) -> Tensor: - return _op_integer(Conv, [x,w], [x_zero_point,w_zero_point], **{"B":B, **opts}) + def ConvInteger(x: Tensor, w: Tensor, x_zero_point:Tensor = Tensor(0), w_zero_point:Tensor = Tensor(0), **opts) -> Tensor: + return _op_integer(Conv, [x,w], [x_zero_point,w_zero_point], **opts) def MatMulInteger(A: Tensor, B: Tensor, a_zero_point: Tensor = Tensor(0), b_zero_point: Tensor = Tensor(0)) -> Tensor: return _op_integer(Tensor.matmul, [A,B], [a_zero_point,b_zero_point])