diff --git a/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/ArgMaxOperationsConverterTest.kt b/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/ArgMaxOperationsConverterTest.kt index 31f1b7edc..e164e3db7 100644 --- a/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/ArgMaxOperationsConverterTest.kt +++ b/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/ArgMaxOperationsConverterTest.kt @@ -1,14 +1,19 @@ package sk.ainet.compile.hlo import sk.ainet.compile.hlo.converters.ArgMaxOperationsConverter +import sk.ainet.lang.dag.argMax +import sk.ainet.lang.dag.dag import sk.ainet.lang.graph.GraphNode +import sk.ainet.lang.graph.dsl.toComputeGraph import sk.ainet.lang.tensor.Tensor import sk.ainet.lang.tensor.ops.Operation import sk.ainet.lang.tensor.ops.TensorSpec import sk.ainet.lang.tensor.ops.ValidationResult import sk.ainet.lang.types.DType +import sk.ainet.lang.types.FP32 import kotlin.test.Test import kotlin.test.assertEquals +import kotlin.test.assertFalse import kotlin.test.assertIs import kotlin.test.assertTrue @@ -57,6 +62,24 @@ class ArgMaxOperationsConverterTest { assertTrue(ops.any { it.contains("-> tensor<1x24xi32>") }, "output is the i32 index tensor") } + @Test + fun testArgMaxThroughDagBuilderProducesCompilableIr() { + // Regression for SKaiNET#876. The dag{} builder path (not just a hand-built node) must infer the + // argMax output spec as a REDUCED, i32 index tensor. Before the fix `inferDagOutputSpecs` had no + // argMax case, so it echoed operand-0's f32 dtype + full shape; the converter then emitted + // `dense : tensor<...xf32>` (an int literal for an f32 constant — iree-compile rejects it) and a + // final reduce that didn't collapse the reduced dim (`1x4x8` instead of `1x4`). + val program = dag { + val x = input("logits", TensorSpec("logits", listOf(1, 4, 8), "FP32")) + output(argMax(x, 2, "am")) + } + val mlir = StableHloConverterFactory.createExtended().convert(program.toComputeGraph(), "argmax_dag").content + + assertTrue(mlir.contains("dense<8> : tensor<1x4x8xi32>"), "sentinel must be an i32 tensor:\n$mlir") + assertFalse(mlir.contains("dense<8> : tensor<1x4x8xf32>"), "no int-literal f32 constant (the #876 bug):\n$mlir") + assertTrue(mlir.contains("-> tensor<1x4xi32>"), "final reduce must collapse to the reduced i32 shape:\n$mlir") + } + @Test fun testInvalidOperandCount() { val operation = createMockOperation("argMax", mapOf("dim" to 1)) diff --git a/skainet-lang/skainet-lang-dag/src/commonMain/kotlin/sk/ainet/lang/dag/GraphDsl.kt b/skainet-lang/skainet-lang-dag/src/commonMain/kotlin/sk/ainet/lang/dag/GraphDsl.kt index 34eb9cc47..9dc5f5216 100644 --- a/skainet-lang/skainet-lang-dag/src/commonMain/kotlin/sk/ainet/lang/dag/GraphDsl.kt +++ b/skainet-lang/skainet-lang-dag/src/commonMain/kotlin/sk/ainet/lang/dag/GraphDsl.kt @@ -173,11 +173,11 @@ public class DagBuilder { nodeId: String ): List? { val input = inputs.firstOrNull()?.spec - fun spec(shape: List?): List = listOf( + fun spec(shape: List?, dtype: String = input?.dtype ?: "unknown"): List = listOf( TensorSpec( name = "${nodeId}_out0", shape = shape, - dtype = input?.dtype ?: "unknown", + dtype = dtype, requiresGrad = input?.requiresGrad ?: false, ), ) @@ -192,6 +192,17 @@ public class DagBuilder { input ?: return null return spec(reductionOutputShape(input.shape, operation.parameters["dim"] as? Int ?: operation.parameters["axis"] as? Int)) } + "argmax", "argmin" -> { + // Index reduction: removes the reduced dim (like sum/mean) AND changes the dtype to an + // integer index (i32) — NOT the input's float dtype. Without this the node echoes operand-0's + // f32, and the StableHloConverter then emits an int literal for an f32 constant + a final + // reduce that doesn't collapse the reduced dim, producing IR iree-compile rejects. (SKaiNET#876) + input ?: return null + return spec( + reductionOutputShape(input.shape, operation.parameters["dim"] as? Int ?: operation.parameters["axis"] as? Int), + dtype = "Int32", + ) + } "reshape", "view" -> { val target = reshapeTargetShape(operation) ?: return null return spec(target)