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")) }