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
4 changes: 4 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,11 @@ public class ScalarOperationsConverter : StableHloOperationConverter {

override val supportedOperations: Set<String> = 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(
Expand All @@ -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}",
Expand Down Expand Up @@ -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<String>,
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<?xf32>"

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.
Expand Down
Original file line number Diff line number Diff line change
@@ -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<HloConversionException> {
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<String, Any>): Operation = object : Operation {
override val name: String = opName
override val type: String = opType
override val parameters: Map<String, Any> = params
override fun <T : DType, V> execute(inputs: List<Tensor<T, V>>): List<Tensor<T, V>> =
throw UnsupportedOperationException("test fixture only")
override fun validateInputs(inputs: List<TensorSpec>): ValidationResult = ValidationResult.Valid
override fun inferOutputs(inputs: List<TensorSpec>): List<TensorSpec> = emptyList()
override fun clone(newParameters: Map<String, Any>): Operation = this
override fun serialize(): Map<String, Any> = mapOf("name" to name, "type" to type)
}
}
Loading