Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
@@ -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

Expand Down Expand Up @@ -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<N> : 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<FP32>("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))
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -173,11 +173,11 @@ public class DagBuilder {
nodeId: String
): List<TensorSpec>? {
val input = inputs.firstOrNull()?.spec
fun spec(shape: List<Int>?): List<TensorSpec> = listOf(
fun spec(shape: List<Int>?, dtype: String = input?.dtype ?: "unknown"): List<TensorSpec> = listOf(
TensorSpec(
name = "${nodeId}_out0",
shape = shape,
dtype = input?.dtype ?: "unknown",
dtype = dtype,
requiresGrad = input?.requiresGrad ?: false,
),
)
Expand All @@ -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)
Expand Down
Loading