diff --git a/CHANGELOG.md b/CHANGELOG.md index b8867d62..b41e3c46 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -36,6 +36,10 @@ ### Fixed +- **`clamp` lowers to StableHLO** + ([#1247](https://github.com/SKaiNET-developers/SKaiNET/issues/1247)): the traced + `clamp(x, minVal, maxVal)` op had no converter at all — the first gap the strict gemma3n export + surfaced once failures stopped being comments. Lowers to `stablehlo.clamp` with splat bounds. - **`indexSelect` is now routable in the StableHLO gather converter** ([#1247](https://github.com/SKaiNET-developers/SKaiNET/issues/1247)): the KSP tracing wrapper emits `indexSelect`, but the registry only knew `index_select`, so traced index-select nodes diff --git a/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/converters/ScalarOperationsConverter.kt b/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/converters/ScalarOperationsConverter.kt index e38eca2d..8215af97 100644 --- a/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/converters/ScalarOperationsConverter.kt +++ b/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/converters/ScalarOperationsConverter.kt @@ -20,7 +20,11 @@ public class ScalarOperationsConverter : StableHloOperationConverter { override val supportedOperations: Set = setOf( "addScalar", "subScalar", "mulScalar", "divScalar", - "rsubScalar", "rdivScalar" + "rsubScalar", "rdivScalar", + // clamp(tensor, minVal, maxVal): two scalar attributes, one tensor + // operand — the first registry gap the strict gemma3n export + // surfaced (#1247). Lowers to stablehlo.clamp with splat bounds. + "clamp" ) override fun convert( @@ -35,6 +39,10 @@ public class ScalarOperationsConverter : StableHloOperationConverter { ) } + if (node.operation.name == "clamp") { + return convertClamp(node, operands, context) + } + val scalar = extractScalar(node.operation.parameters) ?: return ConversionResult.Failure( "${node.operation.name} scalar parameter missing on node ${node.id}", @@ -76,6 +84,45 @@ public class ScalarOperationsConverter : StableHloOperationConverter { ) } + /** + * `clamp(x, minVal, maxVal)` → `stablehlo.clamp %min, %x, %max`, with the + * bounds materialized as splat constants of the output type. The tracer + * records the bounds as `minVal` / `maxVal` (see `TensorOps.clamp`). + */ + private fun convertClamp( + node: GraphNode, + operands: List, + context: ConversionContext + ): ConversionResult { + val params = node.operation.parameters + val minVal = (params["minVal"] as? Number)?.let { formatFloat(it.toDouble()) } + val maxVal = (params["maxVal"] as? Number)?.let { formatFloat(it.toDouble()) } + if (minVal == null || maxVal == null) { + return ConversionResult.Failure( + "clamp requires minVal/maxVal parameters on node ${node.id}, got ${params.keys}", + "Unsupported clamp (missing bounds) for node ${node.id}" + ) + } + val outputSpec = node.outputs.firstOrNull() + val outputType = outputSpec?.let { context.getTypeMapper().mapTensorType(it) } + ?: "tensor" + + val minValue = context.nextTempValue() + val minOp = "$minValue = stablehlo.constant dense<$minVal> : $outputType" + context.emitOperation(minOp) + val maxValue = context.nextTempValue() + val maxOp = "$maxValue = stablehlo.constant dense<$maxVal> : $outputType" + context.emitOperation(maxOp) + + val resultValue = context.nextTempValue() + val op = "$resultValue = stablehlo.clamp $minValue, ${operands[0]}, $maxValue : $outputType" + context.emitOperation(op) + return ConversionResult.Success( + outputValueName = resultValue, + emittedOperations = listOf(minOp, maxOp, op) + ) + } + /** * The scalar is stored under "b" by the KSP-generated tracer (see * `OpAttributeFactory.scalarOp`). Accept "scalar" as an alias too. diff --git a/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/ClampConverterTest.kt b/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/ClampConverterTest.kt new file mode 100644 index 00000000..1df5f6da --- /dev/null +++ b/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/ClampConverterTest.kt @@ -0,0 +1,84 @@ +package sk.ainet.compile.hlo + +import sk.ainet.lang.graph.DefaultComputeGraph +import sk.ainet.lang.graph.GraphEdge +import sk.ainet.lang.graph.GraphNode +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 kotlin.test.Test +import kotlin.test.assertFailsWith +import kotlin.test.assertFalse +import kotlin.test.assertTrue + +/** + * `clamp(x, minVal, maxVal)` was the first registry gap the strict gemma3n + * export surfaced (#1247): the tracer records it with `minVal`/`maxVal` + * attributes and one tensor operand, and no converter claimed it. + */ +class ClampConverterTest { + + @Test + fun clamp_lowers_to_stablehlo_clamp_with_splat_bounds() { + val graph = DefaultComputeGraph() + val input = GraphNode( + id = "x", + operation = fixtureOp("input", "input", emptyMap()), + inputs = emptyList(), + outputs = listOf(TensorSpec("x", listOf(2, 3), "FP32")) + ) + val clamp = GraphNode( + id = "c1", + operation = fixtureOp("clamp", "trace", mapOf("minVal" to -1.5f, "maxVal" to 2.0f)), + inputs = listOf(TensorSpec("x", listOf(2, 3), "FP32")), + outputs = listOf(TensorSpec("y", listOf(2, 3), "FP32")) + ) + graph.addNode(input) + graph.addNode(clamp) + graph.addEdge(GraphEdge("e1", input, clamp, 0, 0, input.outputs[0])) + + val module = StableHloConverterFactory.createBasic().convert(graph, "clamp_test") + val content = module.content + assertTrue(content.contains("stablehlo.constant dense<-1.5> : tensor<2x3xf32>"), content) + assertTrue(content.contains("stablehlo.constant dense<2.0> : tensor<2x3xf32>"), content) + assertTrue( + Regex("""stablehlo\.clamp %v\d+, %arg0, %v\d+ : tensor<2x3xf32>""").containsMatchIn(content), + "clamp must emit `stablehlo.clamp %min, %x, %max`:\n$content" + ) + assertFalse(content.contains("No converter found"), content) + assertFalse(content.contains("Conversion failed"), content) + } + + @Test + fun clamp_without_bounds_is_a_failure_not_a_registry_miss() { + val graph = DefaultComputeGraph() + val input = GraphNode("x", fixtureOp("input", "input", emptyMap()), emptyList(), listOf(TensorSpec("x", listOf(2), "FP32"))) + val clamp = GraphNode("c1", fixtureOp("clamp", "trace", emptyMap()), listOf(TensorSpec("x", listOf(2), "FP32")), listOf(TensorSpec("y", listOf(2), "FP32"))) + graph.addNode(input); graph.addNode(clamp) + graph.addEdge(GraphEdge("e1", input, clamp, 0, 0, input.outputs[0])) + + // Under the default STRICT policy (#1248) a converter Failure aborts the + // conversion; the point here is that it is a named Failure with the + // bounds diagnostic, not a "No converter found" registry miss. + val e = assertFailsWith { + StableHloConverterFactory.createBasic().convert(graph, "clamp_missing") + } + val message = e.message ?: "" + assertTrue("clamp requires minVal/maxVal" in message, message) + assertFalse("No converter found" in message, message) + } + + private fun fixtureOp(opName: String, opType: String, params: Map): Operation = object : Operation { + override val name: String = opName + override val type: String = opType + override val parameters: Map = params + override fun execute(inputs: List>): List> = + throw UnsupportedOperationException("test fixture only") + override fun validateInputs(inputs: List): ValidationResult = ValidationResult.Valid + override fun inferOutputs(inputs: List): List = emptyList() + override fun clone(newParameters: Map): Operation = this + override fun serialize(): Map = mapOf("name" to name, "type" to type) + } +}