From 7d374c7e8223540aa5ad82bccc22d5bb246d12fa Mon Sep 17 00:00:00 2001 From: Michal Harakal Date: Tue, 23 Jun 2026 10:19:15 +0200 Subject: [PATCH] fix(hlo): GroupNorm emits real stablehlo.reduce (IREE-compilable) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The 0.32.0 GroupNorm converter lowered mean/variance with `stablehlo.custom_call @reduce_mean` / `@reduce_variance` (mirroring the LayerNorm style). Those custom_calls are opaque to stock `iree-compile`, so a groupNorm module exported fine but failed to compile on IREE — caught by the downstream skainet-iree-conformance harness (groupnorm:compile). Emit real `stablehlo.reduce` (add region) + divide instead, computing variance as E[x^2] - E[x]^2 (population, ddof=0) — exactly as the standalone `sum`/`mean`/`variance` converters already do. The module now compiles on llvm-cpu and runs. Verified end-to-end via the conformance harness against this source build: iree-compile + iree-run-module + numpy validate -> PASS, max_abs_err=1.2e-7. GroupNormConverterTest updated to assert real `stablehlo.reduce` and the absence of any `custom_call`. (LayerNorm/RMSNorm still use the @reduce_* custom_calls; making them IREE-compilable the same way is a follow-up.) Co-Authored-By: Claude Opus 4.8 (1M context) --- .../NeuralNetOperationsConverter.kt | 28 +++++++++++++++---- .../compile/hlo/GroupNormConverterTest.kt | 13 +++++---- 2 files changed, 30 insertions(+), 11 deletions(-) diff --git a/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/converters/NeuralNetOperationsConverter.kt b/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/converters/NeuralNetOperationsConverter.kt index ca6165a06..bd87db0c3 100644 --- a/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/converters/NeuralNetOperationsConverter.kt +++ b/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/converters/NeuralNetOperationsConverter.kt @@ -509,9 +509,16 @@ public class NeuralNetOperationsConverter : StableHloOperationConverter { val offsetOperand: String? = if (operands.size > 2) operands[2] else null val grouped = context.nextTempValue() + val zeroInit = context.nextTempValue() + val countConst = context.nextTempValue() + val sumX = context.nextTempValue() val meanValue = context.nextTempValue() val meanBroadcast = context.nextTempValue() val centered = context.nextTempValue() + val squared = context.nextTempValue() + val sumSq = context.nextTempValue() + val meanSq = context.nextTempValue() + val meanSquared = context.nextTempValue() val varValue = context.nextTempValue() val epsConst = context.nextTempValue() val epsBroadcast = context.nextTempValue() @@ -527,16 +534,25 @@ public class NeuralNetOperationsConverter : StableHloOperationConverter { // spatial into one trailing axis so a single-axis reduction is per-group. operations += "$grouped = stablehlo.reshape $xInput : ($outputType) -> $groupedType" - // mean(xg) over the trailing axis, broadcast back, mean-center. - operations += "$meanValue = stablehlo.custom_call @reduce_mean($grouped) " + - "{dimensions = [2], keepdim = false} : $reducedType" + // Reductions use real `stablehlo.reduce` (not @reduce_* custom_calls) so the module + // compiles on stock IREE. mean(xg) = sum(xg) / M over the trailing axis; broadcast + // back and mean-center. + operations += "$zeroInit = stablehlo.constant dense<0.0> : tensor<$elementType>" + operations += "$countConst = stablehlo.constant dense<${perGroup}.0> : $reducedType" + operations += "$sumX = stablehlo.reduce($grouped init: $zeroInit) " + + "applies stablehlo.add across dimensions = [2] : ($groupedType, tensor<$elementType>) -> $reducedType" + operations += "$meanValue = stablehlo.divide $sumX, $countConst : $reducedType" operations += "$meanBroadcast = stablehlo.broadcast_in_dim $meanValue, " + "dims = [0, 1] : ($reducedType) -> $groupedType" operations += "$centered = stablehlo.subtract $grouped, $meanBroadcast : $groupedType" - // var(xg) over the trailing axis; std = sqrt(var + eps). - operations += "$varValue = stablehlo.custom_call @reduce_variance($grouped) " + - "{dimensions = [2], keepdim = false} : $reducedType" + // var(xg) = E[xg^2] - E[xg]^2 (population, ddof=0); std = sqrt(var + eps). + operations += "$squared = stablehlo.multiply $grouped, $grouped : $groupedType" + operations += "$sumSq = stablehlo.reduce($squared init: $zeroInit) " + + "applies stablehlo.add across dimensions = [2] : ($groupedType, tensor<$elementType>) -> $reducedType" + operations += "$meanSq = stablehlo.divide $sumSq, $countConst : $reducedType" + operations += "$meanSquared = stablehlo.multiply $meanValue, $meanValue : $reducedType" + operations += "$varValue = stablehlo.subtract $meanSq, $meanSquared : $reducedType" operations += "$epsConst = stablehlo.constant dense<$epsilon> : tensor<$elementType>" operations += "$epsBroadcast = stablehlo.broadcast_in_dim $epsConst, " + "dims = [] : (tensor<$elementType>) -> $reducedType" diff --git a/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/GroupNormConverterTest.kt b/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/GroupNormConverterTest.kt index 8372f3534..740d4c186 100644 --- a/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/GroupNormConverterTest.kt +++ b/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/GroupNormConverterTest.kt @@ -51,9 +51,12 @@ class GroupNormConverterTest { val mlir = module.content assertTrue(mlir.contains("stablehlo.reshape"), "groupNorm must reshape to group the channels and back") - assertTrue(mlir.contains("@reduce_mean"), "groupNorm must lower mean(x) to a real reduction") - assertTrue(mlir.contains("@reduce_variance"), "groupNorm must lower var(x) to a real reduction") - assertTrue(mlir.contains("stablehlo.subtract"), "groupNorm must mean-center") + assertTrue(mlir.contains("stablehlo.reduce("), "groupNorm must lower mean/var to real stablehlo.reduce") + assertFalse( + mlir.contains("custom_call"), + "groupNorm must use real stablehlo.reduce, not @reduce_* custom_calls (those don't compile on IREE)" + ) + assertTrue(mlir.contains("stablehlo.subtract"), "groupNorm must mean-center and form E[x^2]-E[x]^2") assertTrue(mlir.contains("stablehlo.sqrt"), "groupNorm must take sqrt(var + eps)") assertTrue(mlir.contains("stablehlo.divide"), "groupNorm must divide by the std") assertTrue(mlir.contains("stablehlo.broadcast_in_dim"), "groupNorm must broadcast reduced stats back") @@ -68,9 +71,9 @@ class GroupNormConverterTest { val mlir = module.content assertFalse(mlir.contains("@group_norm")) + assertFalse(mlir.contains("custom_call")) assertTrue(mlir.contains("stablehlo.reshape")) - assertTrue(mlir.contains("@reduce_mean")) - assertTrue(mlir.contains("@reduce_variance")) + assertTrue(mlir.contains("stablehlo.reduce(")) assertTrue(mlir.contains("stablehlo.sqrt")) assertTrue(mlir.contains("stablehlo.divide")) }