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
Expand Up @@ -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
}
Expand Down Expand Up @@ -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 <init> ()V
public fun <init> (Lsk/ainet/compile/hlo/StableHloConverter;)V
public fun <init> (Lsk/ainet/compile/hlo/StableHloConverter;Ljava/lang/String;)V
public fun <init> (Lsk/ainet/compile/hlo/StableHloConverter;Ljava/lang/String;Ljava/lang/String;)V
public synthetic fun <init> (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 <init> ()V
public fun <init> (Lsk/ainet/compile/hlo/StableHloGraphExportConverter;)V
public fun <init> (Lsk/ainet/compile/hlo/StableHloGraphExportConverter;Lsk/ainet/compile/hlo/StableHloTextWriter;)V
public synthetic fun <init> (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 <init> (Ljava/lang/String;Ljava/lang/String;Ljava/util/List;Ljava/util/List;Ljava/util/Map;Ljava/util/List;)V
public synthetic fun <init> (Ljava/lang/String;Ljava/lang/String;Ljava/util/List;Ljava/util/List;Ljava/util/Map;Ljava/util/List;ILkotlin/jvm/internal/DefaultConstructorMarker;)V
Expand Down Expand Up @@ -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 <init> ()V
public fun <init> (Ljava/lang/String;)V
public fun <init> (Ljava/lang/String;Ljava/lang/String;)V
public synthetic fun <init> (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 <init> ()V
public final fun areTypesCompatible (Lsk/ainet/lang/tensor/ops/TensorSpec;Lsk/ainet/lang/tensor/ops/TensorSpec;)Z
Expand All @@ -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 <init> ()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 <init> ()V
public fun convert (Lsk/ainet/lang/graph/GraphNode;Ljava/util/List;Lsk/ainet/compile/hlo/ConversionContext;)Lsk/ainet/compile/hlo/ConversionResult;
Expand Down
Original file line number Diff line number Diff line change
@@ -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<GraphExportComponentRole, String> = 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<ComputeGraph, StableHloModule> {

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<StableHloModule, String> {

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<StableHloModule> {
return exportModule(graph, GraphExportContext(backendName = backendName))
}

public fun exportModule(
graph: ComputeGraph,
context: GraphExportContext
): GraphExportResult<StableHloModule> {
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<String> {
return exportText(graph, GraphExportContext(backendName = backendName))
}

public fun exportText(
graph: ComputeGraph,
context: GraphExportContext
): GraphExportResult<String> {
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 <T> stableHloFailureResult(
backendName: String,
stage: GraphExportStage,
exception: Exception,
context: GraphExportContext
): GraphExportResult<T> {
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
)
}
Original file line number Diff line number Diff line change
@@ -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<DType, Any>(),
inputs = emptyList(),
outputs = listOf(TensorSpec("a", listOf(2, 3), "FP32"))
)
val inputB = GraphNode(
id = "b",
operation = InputOperation<DType, Any>(),
inputs = emptyList(),
outputs = listOf(TensorSpec("b", listOf(2, 3), "FP32"))
)
val add = GraphNode(
id = "add1",
operation = AddOperation<DType, Any>(),
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
}
}
Loading