diff --git a/skainet-compile/skainet-compile-hlo/api/jvm/skainet-compile-hlo.api b/skainet-compile/skainet-compile-hlo/api/jvm/skainet-compile-hlo.api index 8277d9594..62483e102 100644 --- a/skainet-compile/skainet-compile-hlo/api/jvm/skainet-compile-hlo.api +++ b/skainet-compile/skainet-compile-hlo/api/jvm/skainet-compile-hlo.api @@ -59,10 +59,13 @@ public final class sk/ainet/compile/hlo/ConversionContext { public final fun getModuleDeclarations ()Ljava/lang/String; public final fun getTypeMapper ()Lsk/ainet/compile/hlo/TypeMapper; public final fun getValueName (Ljava/lang/String;)Ljava/lang/String; + public final fun getValueName (Ljava/lang/String;I)Ljava/lang/String; public final fun getValueType (Ljava/lang/String;)Ljava/lang/String; public final fun nextTempValue ()Ljava/lang/String; public final fun registerExternalParameter (Lsk/ainet/compile/hlo/ExternalParameterRef;)V + public final fun resolveOperands (Lsk/ainet/lang/graph/GraphNode;)Ljava/util/List; public final fun setGraph (Lsk/ainet/lang/graph/ComputeGraph;)V + public final fun setValueName (Ljava/lang/String;ILjava/lang/String;)V public final fun setValueName (Ljava/lang/String;Ljava/lang/String;)V public final fun setValueType (Ljava/lang/String;Ljava/lang/String;)V } @@ -201,6 +204,38 @@ public final class sk/ainet/compile/hlo/StableHloConverterFactory { public static synthetic fun createFast$default (Lsk/ainet/compile/hlo/ConstantMaterializationPolicy;ILjava/lang/Object;)Lsk/ainet/compile/hlo/StableHloConverter; } +public final class sk/ainet/compile/hlo/StableHloExportArchitecture { + public static final field INSTANCE Lsk/ainet/compile/hlo/StableHloExportArchitecture; + public static final field backendName Ljava/lang/String; + public final fun getComponentNames ()Ljava/util/Map; +} + +public final class sk/ainet/compile/hlo/StableHloGraphExportConverter : sk/ainet/compile/export/GraphExportConverter { + public fun ()V + public fun (Lsk/ainet/compile/hlo/StableHloConverter;)V + public fun (Lsk/ainet/compile/hlo/StableHloConverter;Ljava/lang/String;)V + public fun (Lsk/ainet/compile/hlo/StableHloConverter;Ljava/lang/String;Ljava/lang/String;)V + public synthetic fun (Lsk/ainet/compile/hlo/StableHloConverter;Ljava/lang/String;Ljava/lang/String;ILkotlin/jvm/internal/DefaultConstructorMarker;)V + public synthetic fun convert (Ljava/lang/Object;Lsk/ainet/compile/export/GraphExportContext;)Ljava/lang/Object; + public fun convert (Lsk/ainet/lang/graph/ComputeGraph;Lsk/ainet/compile/export/GraphExportContext;)Lsk/ainet/compile/hlo/StableHloModule; + public fun getBackendName ()Ljava/lang/String; + public final fun getFunctionName ()Ljava/lang/String; +} + +public final class sk/ainet/compile/hlo/StableHloGraphExporter { + public fun ()V + public fun (Lsk/ainet/compile/hlo/StableHloGraphExportConverter;)V + public fun (Lsk/ainet/compile/hlo/StableHloGraphExportConverter;Lsk/ainet/compile/hlo/StableHloTextWriter;)V + public synthetic fun (Lsk/ainet/compile/hlo/StableHloGraphExportConverter;Lsk/ainet/compile/hlo/StableHloTextWriter;ILkotlin/jvm/internal/DefaultConstructorMarker;)V + public final fun exportModule (Lsk/ainet/lang/graph/ComputeGraph;)Lsk/ainet/compile/export/GraphExportResult; + public final fun exportModule (Lsk/ainet/lang/graph/ComputeGraph;Lsk/ainet/compile/export/GraphExportContext;)Lsk/ainet/compile/export/GraphExportResult; + public final fun exportText (Lsk/ainet/lang/graph/ComputeGraph;)Lsk/ainet/compile/export/GraphExportResult; + public final fun exportText (Lsk/ainet/lang/graph/ComputeGraph;Lsk/ainet/compile/export/GraphExportContext;)Lsk/ainet/compile/export/GraphExportResult; + public final fun getBackendName ()Ljava/lang/String; + public final fun getConverter ()Lsk/ainet/compile/hlo/StableHloGraphExportConverter; + public final fun getWriter ()Lsk/ainet/compile/hlo/StableHloTextWriter; +} + public final class sk/ainet/compile/hlo/StableHloModule { public fun (Ljava/lang/String;Ljava/lang/String;Ljava/util/List;Ljava/util/List;Ljava/util/Map;Ljava/util/List;)V public synthetic fun (Ljava/lang/String;Ljava/lang/String;Ljava/util/List;Ljava/util/List;Ljava/util/Map;Ljava/util/List;ILkotlin/jvm/internal/DefaultConstructorMarker;)V @@ -259,6 +294,17 @@ public final class sk/ainet/compile/hlo/StableHloOptimizer$Companion { public final fun createDefault ()Lsk/ainet/compile/hlo/StableHloOptimizer; } +public final class sk/ainet/compile/hlo/StableHloTextWriter : sk/ainet/compile/export/GraphExportWriter { + public fun ()V + public fun (Ljava/lang/String;)V + public fun (Ljava/lang/String;Ljava/lang/String;)V + public synthetic fun (Ljava/lang/String;Ljava/lang/String;ILkotlin/jvm/internal/DefaultConstructorMarker;)V + public fun getBackendName ()Ljava/lang/String; + public final fun getLogicalPath ()Ljava/lang/String; + public synthetic fun write (Ljava/lang/Object;Lsk/ainet/compile/export/GraphExportContext;)Ljava/lang/Object; + public fun write (Lsk/ainet/compile/hlo/StableHloModule;Lsk/ainet/compile/export/GraphExportContext;)Ljava/lang/String; +} + public final class sk/ainet/compile/hlo/TypeMapper { public fun ()V public final fun areTypesCompatible (Lsk/ainet/lang/tensor/ops/TensorSpec;Lsk/ainet/lang/tensor/ops/TensorSpec;)Z @@ -277,6 +323,12 @@ public final class sk/ainet/compile/hlo/converters/ActivationOperationsConverter public fun getSupportedOperations ()Ljava/util/Set; } +public final class sk/ainet/compile/hlo/converters/AttentionOperationsConverter : sk/ainet/compile/hlo/StableHloOperationConverter { + public fun ()V + public fun convert (Lsk/ainet/lang/graph/GraphNode;Ljava/util/List;Lsk/ainet/compile/hlo/ConversionContext;)Lsk/ainet/compile/hlo/ConversionResult; + public fun getSupportedOperations ()Ljava/util/Set; +} + public final class sk/ainet/compile/hlo/converters/BasicMathConverter : sk/ainet/compile/hlo/StableHloOperationConverter { public fun ()V public fun convert (Lsk/ainet/lang/graph/GraphNode;Ljava/util/List;Lsk/ainet/compile/hlo/ConversionContext;)Lsk/ainet/compile/hlo/ConversionResult; diff --git a/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/StableHloGraphExport.kt b/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/StableHloGraphExport.kt new file mode 100644 index 000000000..a2e852d91 --- /dev/null +++ b/skainet-compile/skainet-compile-hlo/src/commonMain/kotlin/sk/ainet/compile/hlo/StableHloGraphExport.kt @@ -0,0 +1,187 @@ +package sk.ainet.compile.hlo + +import sk.ainet.compile.export.GraphExportArtifact +import sk.ainet.compile.export.GraphExportArtifactRole +import sk.ainet.compile.export.GraphExportComponentRole +import sk.ainet.compile.export.GraphExportContext +import sk.ainet.compile.export.GraphExportConverter +import sk.ainet.compile.export.GraphExportResult +import sk.ainet.compile.export.GraphExportStage +import sk.ainet.compile.export.GraphExportStatus +import sk.ainet.compile.export.GraphExportWriter +import sk.ainet.lang.graph.ComputeGraph + +/** + * StableHLO component mapping in the shared graph-export architecture. + */ +public object StableHloExportArchitecture { + public const val backendName: String = "stablehlo" + + public val componentNames: Map = mapOf( + GraphExportComponentRole.CONVERTER to "StableHloConverter", + GraphExportComponentRole.CONTEXT to "ConversionContext", + GraphExportComponentRole.REGISTRY to "StableHloOperationRegistry", + GraphExportComponentRole.FACTORY to "StableHloConverterFactory", + GraphExportComponentRole.WRITER to "StableHloTextWriter", + GraphExportComponentRole.VERIFIER to "MlirValidator" + ) +} + +/** + * Adapter that exposes [StableHloConverter] through the shared export contract. + */ +public class StableHloGraphExportConverter @kotlin.jvm.JvmOverloads constructor( + private val converter: StableHloConverter = StableHloConverterFactory.createBasic(), + public val functionName: String = "main", + override val backendName: String = StableHloExportArchitecture.backendName +) : GraphExportConverter { + + override fun convert(input: ComputeGraph, context: GraphExportContext): StableHloModule { + val resolvedFunctionName = context.targetName ?: functionName + context.info( + stage = GraphExportStage.LOWERING, + code = "stablehlo.lowering.started", + message = "Lowering ComputeGraph to StableHLO MLIR.", + details = mapOf("functionName" to resolvedFunctionName) + ) + + val module = converter.convert(input, resolvedFunctionName) + + context.info( + stage = GraphExportStage.LOWERING, + code = "stablehlo.lowering.completed", + message = "Lowered ComputeGraph to StableHLO MLIR.", + details = mapOf( + "functionName" to module.functionName, + "inputs" to module.inputSpecs.size.toString(), + "outputs" to module.outputSpecs.size.toString(), + "externalParameters" to module.externalParameters.size.toString() + ) + ) + return module + } +} + +/** + * Shared-contract writer that renders a [StableHloModule] as MLIR text. + */ +public class StableHloTextWriter @kotlin.jvm.JvmOverloads constructor( + public val logicalPath: String? = null, + override val backendName: String = StableHloExportArchitecture.backendName +) : GraphExportWriter { + + override fun write(intermediate: StableHloModule, context: GraphExportContext): String { + val artifactPath = logicalPath ?: "${intermediate.functionName}.stablehlo.mlir" + context.addArtifact( + GraphExportArtifact( + path = artifactPath, + role = GraphExportArtifactRole.SOURCE, + description = "StableHLO MLIR module text", + metadata = mapOf( + "functionName" to intermediate.functionName, + "format" to "mlir" + ) + ) + ) + context.info( + stage = GraphExportStage.WRITING, + code = "stablehlo.writing.text", + message = "Rendered StableHLO module as MLIR text.", + details = mapOf( + "path" to artifactPath, + "characters" to intermediate.content.length.toString() + ) + ) + return intermediate.content + } +} + +/** + * Convenience facade that composes StableHLO lowering and writing into shared results. + */ +public class StableHloGraphExporter @kotlin.jvm.JvmOverloads constructor( + public val converter: StableHloGraphExportConverter = StableHloGraphExportConverter(), + public val writer: StableHloTextWriter = StableHloTextWriter() +) { + public val backendName: String + get() = converter.backendName + + public fun exportModule(graph: ComputeGraph): GraphExportResult { + return exportModule(graph, GraphExportContext(backendName = backendName)) + } + + public fun exportModule( + graph: ComputeGraph, + context: GraphExportContext + ): GraphExportResult { + return try { + val module = converter.convert(graph, context) + GraphExportResult.success( + backendName = backendName, + output = module, + diagnostics = context.diagnosticReport(), + artifacts = context.artifacts, + metadata = context.metadata + ) + } catch (exception: Exception) { + stableHloFailureResult( + backendName = backendName, + stage = GraphExportStage.LOWERING, + exception = exception, + context = context + ) + } + } + + public fun exportText(graph: ComputeGraph): GraphExportResult { + return exportText(graph, GraphExportContext(backendName = backendName)) + } + + public fun exportText( + graph: ComputeGraph, + context: GraphExportContext + ): GraphExportResult { + var stage = GraphExportStage.LOWERING + return try { + val module = converter.convert(graph, context) + stage = GraphExportStage.WRITING + val text = writer.write(module, context) + GraphExportResult.success( + backendName = backendName, + output = text, + diagnostics = context.diagnosticReport(), + artifacts = context.artifacts, + metadata = context.metadata + ) + } catch (exception: Exception) { + stableHloFailureResult( + backendName = backendName, + stage = stage, + exception = exception, + context = context + ) + } + } +} + +private fun stableHloFailureResult( + backendName: String, + stage: GraphExportStage, + exception: Exception, + context: GraphExportContext +): GraphExportResult { + val reason = exception.message ?: exception.toString() + context.error( + stage = stage, + code = "stablehlo.export.failed", + message = "StableHLO export failed: $reason" + ) + return GraphExportResult( + backendName = backendName, + status = GraphExportStatus.FAILED, + output = null, + diagnostics = context.diagnosticReport(), + artifacts = context.artifacts, + metadata = context.metadata + ) +} diff --git a/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/StableHloGraphExportTest.kt b/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/StableHloGraphExportTest.kt new file mode 100644 index 000000000..cac1cd0b4 --- /dev/null +++ b/skainet-compile/skainet-compile-hlo/src/commonTest/kotlin/sk/ainet/compile/hlo/StableHloGraphExportTest.kt @@ -0,0 +1,116 @@ +package sk.ainet.compile.hlo + +import kotlin.test.Test +import kotlin.test.assertEquals +import kotlin.test.assertFalse +import kotlin.test.assertTrue +import sk.ainet.compile.export.GraphExportArtifactRole +import sk.ainet.compile.export.GraphExportComponentRole +import sk.ainet.compile.export.GraphExportContext +import sk.ainet.compile.export.GraphExportStage +import sk.ainet.compile.export.GraphExportStatus +import sk.ainet.lang.graph.DefaultComputeGraph +import sk.ainet.lang.graph.GraphEdge +import sk.ainet.lang.graph.GraphNode +import sk.ainet.lang.tensor.ops.AddOperation +import sk.ainet.lang.tensor.ops.InputOperation +import sk.ainet.lang.tensor.ops.TensorSpec +import sk.ainet.lang.types.DType + +class StableHloGraphExportTest { + + @Test + fun graphExportConverterPreservesStableHloOutput() { + val graph = simpleAddGraph() + val directConverter = StableHloConverterFactory.createBasic() + val expected = directConverter.convert(graph, "shared_export") + val context = GraphExportContext( + backendName = StableHloExportArchitecture.backendName, + targetName = "shared_export" + ) + + val actual = StableHloGraphExportConverter(converter = directConverter).convert(graph, context) + + assertEquals(expected.content, actual.content) + assertEquals(expected.functionName, actual.functionName) + assertFalse(context.diagnosticReport().hasErrors) + assertTrue(context.diagnostics.any { it.code == "stablehlo.lowering.started" }) + assertTrue(context.diagnostics.any { it.code == "stablehlo.lowering.completed" }) + } + + @Test + fun textWriterRecordsLogicalStableHloArtifact() { + val module = toStableHlo(simpleAddGraph(), "text_export") + val context = GraphExportContext(backendName = StableHloExportArchitecture.backendName) + val writer = StableHloTextWriter(logicalPath = "build/generated/text_export.mlir") + + val text = writer.write(module, context) + + assertEquals(module.content, text) + assertEquals(1, context.artifacts.size) + assertEquals("build/generated/text_export.mlir", context.artifacts.single().path) + assertEquals(GraphExportArtifactRole.SOURCE, context.artifacts.single().role) + assertTrue(context.diagnostics.any { it.stage == GraphExportStage.WRITING }) + } + + @Test + fun graphExporterReturnsSharedResultEnvelopeForStableHloText() { + val graph = simpleAddGraph() + val expected = StableHloConverterFactory.createBasic().convert(graph, "result_export") + val context = GraphExportContext( + backendName = StableHloExportArchitecture.backendName, + targetName = "result_export" + ) + val exporter = StableHloGraphExporter( + writer = StableHloTextWriter(logicalPath = "result_export.mlir") + ) + + val result = exporter.exportText(graph, context) + + assertEquals(GraphExportStatus.SUCCESS, result.status) + assertEquals(expected.content, result.requireSuccess()) + assertFalse(result.diagnostics.hasErrors) + assertEquals(1, result.artifacts.size) + assertEquals("result_export.mlir", result.artifacts.single().path) + assertEquals( + "StableHloConverter", + StableHloExportArchitecture.componentNames[GraphExportComponentRole.CONVERTER] + ) + assertEquals( + "StableHloTextWriter", + StableHloExportArchitecture.componentNames[GraphExportComponentRole.WRITER] + ) + } + + private fun simpleAddGraph(): DefaultComputeGraph { + val graph = DefaultComputeGraph() + val inputA = GraphNode( + id = "a", + operation = InputOperation(), + inputs = emptyList(), + outputs = listOf(TensorSpec("a", listOf(2, 3), "FP32")) + ) + val inputB = GraphNode( + id = "b", + operation = InputOperation(), + inputs = emptyList(), + outputs = listOf(TensorSpec("b", listOf(2, 3), "FP32")) + ) + val add = GraphNode( + id = "add1", + operation = AddOperation(), + inputs = listOf( + TensorSpec("a", listOf(2, 3), "FP32"), + TensorSpec("b", listOf(2, 3), "FP32") + ), + outputs = listOf(TensorSpec("c", listOf(2, 3), "FP32")) + ) + + graph.addNode(inputA) + graph.addNode(inputB) + graph.addNode(add) + graph.addEdge(GraphEdge("e1", inputA, add, 0, 0, inputA.outputs[0])) + graph.addEdge(GraphEdge("e2", inputB, add, 0, 1, inputB.outputs[0])) + return graph + } +}