From 93dc8b8470bb156083cfb62f28db1a264384eb2f Mon Sep 17 00:00:00 2001 From: Michal Harakal Date: Sat, 6 Jun 2026 23:40:47 +0200 Subject: [PATCH] feat(minerva): validate phase-one graph compatibility Closes #691 --- .../api/skainet-compile-minerva.api | 99 +++- .../minerva/MinervaCompatibilityValidator.kt | 433 ++++++++++++++++++ .../compile/minerva/MinervaExportFacade.kt | 75 ++- .../compile/minerva/MinervaExportModels.kt | 63 ++- .../MinervaCompatibilityValidatorTest.kt | 121 +++++ .../minerva/MinervaExportFacadeTest.kt | 66 ++- .../compile/minerva/MinervaGraphFixtures.kt | 187 ++++++++ 7 files changed, 963 insertions(+), 81 deletions(-) create mode 100644 skainet-compile/skainet-compile-minerva/src/commonMain/kotlin/sk/ainet/compile/minerva/MinervaCompatibilityValidator.kt create mode 100644 skainet-compile/skainet-compile-minerva/src/commonTest/kotlin/sk/ainet/compile/minerva/MinervaCompatibilityValidatorTest.kt create mode 100644 skainet-compile/skainet-compile-minerva/src/commonTest/kotlin/sk/ainet/compile/minerva/MinervaGraphFixtures.kt diff --git a/skainet-compile/skainet-compile-minerva/api/skainet-compile-minerva.api b/skainet-compile/skainet-compile-minerva/api/skainet-compile-minerva.api index 6fe8d6cc1..f02dd2908 100644 --- a/skainet-compile/skainet-compile-minerva/api/skainet-compile-minerva.api +++ b/skainet-compile/skainet-compile-minerva/api/skainet-compile-minerva.api @@ -1,3 +1,87 @@ +public final class sk/ainet/compile/minerva/MinervaCompatibilityIssue { + public fun (Lsk/ainet/compile/minerva/MinervaCompatibilityIssueKind;Ljava/lang/String;Ljava/lang/String;Ljava/lang/String;Ljava/lang/String;Ljava/lang/String;Ljava/util/Map;)V + public synthetic fun (Lsk/ainet/compile/minerva/MinervaCompatibilityIssueKind;Ljava/lang/String;Ljava/lang/String;Ljava/lang/String;Ljava/lang/String;Ljava/lang/String;Ljava/util/Map;ILkotlin/jvm/internal/DefaultConstructorMarker;)V + public final fun component1 ()Lsk/ainet/compile/minerva/MinervaCompatibilityIssueKind; + public final fun component2 ()Ljava/lang/String; + public final fun component3 ()Ljava/lang/String; + public final fun component4 ()Ljava/lang/String; + public final fun component5 ()Ljava/lang/String; + public final fun component6 ()Ljava/lang/String; + public final fun component7 ()Ljava/util/Map; + public final fun copy (Lsk/ainet/compile/minerva/MinervaCompatibilityIssueKind;Ljava/lang/String;Ljava/lang/String;Ljava/lang/String;Ljava/lang/String;Ljava/lang/String;Ljava/util/Map;)Lsk/ainet/compile/minerva/MinervaCompatibilityIssue; + public static synthetic fun copy$default (Lsk/ainet/compile/minerva/MinervaCompatibilityIssue;Lsk/ainet/compile/minerva/MinervaCompatibilityIssueKind;Ljava/lang/String;Ljava/lang/String;Ljava/lang/String;Ljava/lang/String;Ljava/lang/String;Ljava/util/Map;ILjava/lang/Object;)Lsk/ainet/compile/minerva/MinervaCompatibilityIssue; + public fun equals (Ljava/lang/Object;)Z + public final fun getCode ()Ljava/lang/String; + public final fun getDetails ()Ljava/util/Map; + public final fun getKind ()Lsk/ainet/compile/minerva/MinervaCompatibilityIssueKind; + public final fun getMessage ()Ljava/lang/String; + public final fun getNodeId ()Ljava/lang/String; + public final fun getOperationName ()Ljava/lang/String; + public final fun getRemediation ()Ljava/lang/String; + public fun hashCode ()I + public fun toString ()Ljava/lang/String; +} + +public final class sk/ainet/compile/minerva/MinervaCompatibilityIssueKind : java/lang/Enum { + public static final field GRAPH_VALIDATION Lsk/ainet/compile/minerva/MinervaCompatibilityIssueKind; + public static final field INCOMPATIBLE_ACTIVATION_PLACEMENT Lsk/ainet/compile/minerva/MinervaCompatibilityIssueKind; + public static final field INVALID_SHAPE Lsk/ainet/compile/minerva/MinervaCompatibilityIssueKind; + public static final field MEMORY_BUDGET_EXCEEDED Lsk/ainet/compile/minerva/MinervaCompatibilityIssueKind; + public static final field MISSING_SHAPE Lsk/ainet/compile/minerva/MinervaCompatibilityIssueKind; + public static final field UNSUPPORTED_OPERATION Lsk/ainet/compile/minerva/MinervaCompatibilityIssueKind; + public static final field UNSUPPORTED_QUANTIZATION Lsk/ainet/compile/minerva/MinervaCompatibilityIssueKind; + public static final field UNSUPPORTED_TOPOLOGY Lsk/ainet/compile/minerva/MinervaCompatibilityIssueKind; + public static fun getEntries ()Lkotlin/enums/EnumEntries; + public static fun valueOf (Ljava/lang/String;)Lsk/ainet/compile/minerva/MinervaCompatibilityIssueKind; + public static fun values ()[Lsk/ainet/compile/minerva/MinervaCompatibilityIssueKind; +} + +public final class sk/ainet/compile/minerva/MinervaCompatibilityReport { + public fun (ZLsk/ainet/compile/export/GraphExportDiagnosticReport;Ljava/util/List;Lsk/ainet/compile/minerva/MinervaTarget;Lsk/ainet/compile/minerva/MinervaQuantization;IIILjava/util/Map;)V + public synthetic fun (ZLsk/ainet/compile/export/GraphExportDiagnosticReport;Ljava/util/List;Lsk/ainet/compile/minerva/MinervaTarget;Lsk/ainet/compile/minerva/MinervaQuantization;IIILjava/util/Map;ILkotlin/jvm/internal/DefaultConstructorMarker;)V + public final fun component1 ()Z + public final fun component2 ()Lsk/ainet/compile/export/GraphExportDiagnosticReport; + public final fun component3 ()Ljava/util/List; + public final fun component4 ()Lsk/ainet/compile/minerva/MinervaTarget; + public final fun component5 ()Lsk/ainet/compile/minerva/MinervaQuantization; + public final fun component6 ()I + public final fun component7 ()I + public final fun component8 ()I + public final fun component9 ()Ljava/util/Map; + public final fun copy (ZLsk/ainet/compile/export/GraphExportDiagnosticReport;Ljava/util/List;Lsk/ainet/compile/minerva/MinervaTarget;Lsk/ainet/compile/minerva/MinervaQuantization;IIILjava/util/Map;)Lsk/ainet/compile/minerva/MinervaCompatibilityReport; + public static synthetic fun copy$default (Lsk/ainet/compile/minerva/MinervaCompatibilityReport;ZLsk/ainet/compile/export/GraphExportDiagnosticReport;Ljava/util/List;Lsk/ainet/compile/minerva/MinervaTarget;Lsk/ainet/compile/minerva/MinervaQuantization;IIILjava/util/Map;ILjava/lang/Object;)Lsk/ainet/compile/minerva/MinervaCompatibilityReport; + public fun equals (Ljava/lang/Object;)Z + public final fun getCompatible ()Z + public final fun getDiagnostics ()Lsk/ainet/compile/export/GraphExportDiagnosticReport; + public final fun getEstimatedFlashBytes ()I + public final fun getEstimatedSramBytes ()I + public final fun getFailed ()Z + public final fun getIssues ()Ljava/util/List; + public final fun getLayerCount ()I + public final fun getMetadata ()Ljava/util/Map; + public final fun getQuantization ()Lsk/ainet/compile/minerva/MinervaQuantization; + public final fun getTarget ()Lsk/ainet/compile/minerva/MinervaTarget; + public fun hashCode ()I + public final fun requireCompatible ()Lsk/ainet/compile/minerva/MinervaCompatibilityReport; + public fun toString ()Ljava/lang/String; +} + +public final class sk/ainet/compile/minerva/MinervaCompatibilityValidator { + public static final field Companion Lsk/ainet/compile/minerva/MinervaCompatibilityValidator$Companion; + public fun ()V + public fun (I)V + public synthetic fun (IILkotlin/jvm/internal/DefaultConstructorMarker;)V + public final fun getSramSafetyBytes ()I + public final fun validate (Lsk/ainet/lang/graph/ComputeGraph;Lsk/ainet/compile/minerva/MinervaExportOptions;)Lsk/ainet/compile/minerva/MinervaCompatibilityReport; + public final fun validate (Lsk/ainet/lang/graph/ComputeGraph;Lsk/ainet/compile/minerva/MinervaExportOptions;Lsk/ainet/compile/export/GraphExportContext;)Lsk/ainet/compile/minerva/MinervaCompatibilityReport; +} + +public final class sk/ainet/compile/minerva/MinervaCompatibilityValidator$Companion { + public final fun getActivationOperations ()Ljava/util/Set; + public final fun getActivationProducerOperations ()Ljava/util/Set; + public final fun getLayerOperations ()Ljava/util/Set; +} + public final class sk/ainet/compile/minerva/MinervaExportBackend { public static final field INSTANCE Lsk/ainet/compile/minerva/MinervaExportBackend; public static final field backendName Ljava/lang/String; @@ -29,11 +113,13 @@ public final class sk/ainet/compile/minerva/MinervaExportBundle { public final class sk/ainet/compile/minerva/MinervaExportFacade { public fun ()V public fun (Ljava/lang/String;)V - public synthetic fun (Ljava/lang/String;ILkotlin/jvm/internal/DefaultConstructorMarker;)V + public fun (Ljava/lang/String;Lsk/ainet/compile/minerva/MinervaCompatibilityValidator;)V + public synthetic fun (Ljava/lang/String;Lsk/ainet/compile/minerva/MinervaCompatibilityValidator;ILkotlin/jvm/internal/DefaultConstructorMarker;)V public final fun exportGraph (Lsk/ainet/lang/graph/ComputeGraph;Lsk/ainet/compile/minerva/MinervaExportOptions;)Lsk/ainet/compile/minerva/MinervaExportResult; public final fun exportModel (Ljava/lang/Object;Lkotlin/jvm/functions/Function1;Lsk/ainet/compile/minerva/MinervaExportOptions;)Lsk/ainet/compile/minerva/MinervaExportResult; public final fun exportModel (Ljava/lang/Object;Lsk/ainet/compile/minerva/MinervaExportOptions;)Lsk/ainet/compile/minerva/MinervaExportResult; public final fun getBackendName ()Ljava/lang/String; + public final fun getCompatibilityValidator ()Lsk/ainet/compile/minerva/MinervaCompatibilityValidator; } public final class sk/ainet/compile/minerva/MinervaExportFailure { @@ -57,6 +143,7 @@ public final class sk/ainet/compile/minerva/MinervaExportFailure { } public final class sk/ainet/compile/minerva/MinervaExportFailureKind : java/lang/Enum { + public static final field COMPATIBILITY_VALIDATION_FAILED Lsk/ainet/compile/minerva/MinervaExportFailureKind; public static final field GRAPH_VALIDATION_FAILED Lsk/ainet/compile/minerva/MinervaExportFailureKind; public static final field NOT_IMPLEMENTED Lsk/ainet/compile/minerva/MinervaExportFailureKind; public static final field RECORDING_FAILED Lsk/ainet/compile/minerva/MinervaExportFailureKind; @@ -104,8 +191,8 @@ public final class sk/ainet/compile/minerva/MinervaExportOptions { } public final class sk/ainet/compile/minerva/MinervaExportResult { - public fun (Lsk/ainet/compile/minerva/MinervaExportOptions;Lsk/ainet/compile/export/GraphExportStatus;Lsk/ainet/compile/minerva/MinervaExportBundle;Lsk/ainet/compile/export/GraphExportDiagnosticReport;Ljava/util/List;Lsk/ainet/compile/minerva/MinervaExportFailure;Ljava/util/Map;)V - public synthetic fun (Lsk/ainet/compile/minerva/MinervaExportOptions;Lsk/ainet/compile/export/GraphExportStatus;Lsk/ainet/compile/minerva/MinervaExportBundle;Lsk/ainet/compile/export/GraphExportDiagnosticReport;Ljava/util/List;Lsk/ainet/compile/minerva/MinervaExportFailure;Ljava/util/Map;ILkotlin/jvm/internal/DefaultConstructorMarker;)V + public fun (Lsk/ainet/compile/minerva/MinervaExportOptions;Lsk/ainet/compile/export/GraphExportStatus;Lsk/ainet/compile/minerva/MinervaExportBundle;Lsk/ainet/compile/export/GraphExportDiagnosticReport;Ljava/util/List;Lsk/ainet/compile/minerva/MinervaExportFailure;Ljava/util/Map;Lsk/ainet/compile/minerva/MinervaCompatibilityReport;)V + public synthetic fun (Lsk/ainet/compile/minerva/MinervaExportOptions;Lsk/ainet/compile/export/GraphExportStatus;Lsk/ainet/compile/minerva/MinervaExportBundle;Lsk/ainet/compile/export/GraphExportDiagnosticReport;Ljava/util/List;Lsk/ainet/compile/minerva/MinervaExportFailure;Ljava/util/Map;Lsk/ainet/compile/minerva/MinervaCompatibilityReport;ILkotlin/jvm/internal/DefaultConstructorMarker;)V public final fun component1 ()Lsk/ainet/compile/minerva/MinervaExportOptions; public final fun component2 ()Lsk/ainet/compile/export/GraphExportStatus; public final fun component3 ()Lsk/ainet/compile/minerva/MinervaExportBundle; @@ -113,11 +200,13 @@ public final class sk/ainet/compile/minerva/MinervaExportResult { public final fun component5 ()Ljava/util/List; public final fun component6 ()Lsk/ainet/compile/minerva/MinervaExportFailure; public final fun component7 ()Ljava/util/Map; - public final fun copy (Lsk/ainet/compile/minerva/MinervaExportOptions;Lsk/ainet/compile/export/GraphExportStatus;Lsk/ainet/compile/minerva/MinervaExportBundle;Lsk/ainet/compile/export/GraphExportDiagnosticReport;Ljava/util/List;Lsk/ainet/compile/minerva/MinervaExportFailure;Ljava/util/Map;)Lsk/ainet/compile/minerva/MinervaExportResult; - public static synthetic fun copy$default (Lsk/ainet/compile/minerva/MinervaExportResult;Lsk/ainet/compile/minerva/MinervaExportOptions;Lsk/ainet/compile/export/GraphExportStatus;Lsk/ainet/compile/minerva/MinervaExportBundle;Lsk/ainet/compile/export/GraphExportDiagnosticReport;Ljava/util/List;Lsk/ainet/compile/minerva/MinervaExportFailure;Ljava/util/Map;ILjava/lang/Object;)Lsk/ainet/compile/minerva/MinervaExportResult; + public final fun component8 ()Lsk/ainet/compile/minerva/MinervaCompatibilityReport; + public final fun copy (Lsk/ainet/compile/minerva/MinervaExportOptions;Lsk/ainet/compile/export/GraphExportStatus;Lsk/ainet/compile/minerva/MinervaExportBundle;Lsk/ainet/compile/export/GraphExportDiagnosticReport;Ljava/util/List;Lsk/ainet/compile/minerva/MinervaExportFailure;Ljava/util/Map;Lsk/ainet/compile/minerva/MinervaCompatibilityReport;)Lsk/ainet/compile/minerva/MinervaExportResult; + public static synthetic fun copy$default (Lsk/ainet/compile/minerva/MinervaExportResult;Lsk/ainet/compile/minerva/MinervaExportOptions;Lsk/ainet/compile/export/GraphExportStatus;Lsk/ainet/compile/minerva/MinervaExportBundle;Lsk/ainet/compile/export/GraphExportDiagnosticReport;Ljava/util/List;Lsk/ainet/compile/minerva/MinervaExportFailure;Ljava/util/Map;Lsk/ainet/compile/minerva/MinervaCompatibilityReport;ILjava/lang/Object;)Lsk/ainet/compile/minerva/MinervaExportResult; public fun equals (Ljava/lang/Object;)Z public final fun getArtifacts ()Ljava/util/List; public final fun getBundle ()Lsk/ainet/compile/minerva/MinervaExportBundle; + public final fun getCompatibilityReport ()Lsk/ainet/compile/minerva/MinervaCompatibilityReport; public final fun getDiagnostics ()Lsk/ainet/compile/export/GraphExportDiagnosticReport; public final fun getFailed ()Z public final fun getFailure ()Lsk/ainet/compile/minerva/MinervaExportFailure; diff --git a/skainet-compile/skainet-compile-minerva/src/commonMain/kotlin/sk/ainet/compile/minerva/MinervaCompatibilityValidator.kt b/skainet-compile/skainet-compile-minerva/src/commonMain/kotlin/sk/ainet/compile/minerva/MinervaCompatibilityValidator.kt new file mode 100644 index 000000000..000957662 --- /dev/null +++ b/skainet-compile/skainet-compile-minerva/src/commonMain/kotlin/sk/ainet/compile/minerva/MinervaCompatibilityValidator.kt @@ -0,0 +1,433 @@ +package sk.ainet.compile.minerva + +import sk.ainet.compile.export.GraphExportContext +import sk.ainet.compile.export.GraphExportStage +import sk.ainet.lang.graph.ComputeGraph +import sk.ainet.lang.graph.GraphNode +import sk.ainet.lang.tensor.ops.ValidationResult + +/** + * Validates the phase-one Minerva graph contract before lowering or compiler invocation. + */ +public class MinervaCompatibilityValidator @kotlin.jvm.JvmOverloads constructor( + public val sramSafetyBytes: Int = 0 +) { + public fun validate( + graph: ComputeGraph, + options: MinervaExportOptions + ): MinervaCompatibilityReport { + val context = GraphExportContext( + backendName = MinervaExportBackend.backendName, + targetName = options.projectName, + metadata = options.toMetadata() + ) + return validate(graph, options, context) + } + + public fun validate( + graph: ComputeGraph, + options: MinervaExportOptions, + context: GraphExportContext + ): MinervaCompatibilityReport { + val issues = mutableListOf() + context.info( + stage = GraphExportStage.VALIDATION, + code = "minerva.compatibility.started", + message = "Started Minerva phase-one compatibility validation.", + details = mapOf( + "target" to options.target.compilerId, + "quantization" to options.quantization.compilerId, + "nodes" to graph.nodes.size.toString() + ) + ) + + if (options.quantization != MinervaQuantization.Q8) { + recordIssue( + issues = issues, + context = context, + kind = MinervaCompatibilityIssueKind.UNSUPPORTED_QUANTIZATION, + code = "minerva.compatibility.unsupported_quantization", + message = "Minerva phase one supports Q8 export only.", + remediation = "Use MinervaQuantization.Q8 for phase-one exports.", + details = mapOf("quantization" to options.quantization.compilerId) + ) + } + + if (graph.nodes.isEmpty()) { + recordIssue( + issues = issues, + context = context, + kind = MinervaCompatibilityIssueKind.GRAPH_VALIDATION, + code = "minerva.compatibility.empty_graph", + message = "Minerva export requires at least one graph node.", + remediation = "Provide a traced or constructed ComputeGraph with input and layer nodes." + ) + return report(options, context, issues, layerCount = 0, estimatedSramBytes = 0, estimatedFlashBytes = 0) + } + + when (val validation = graph.validate()) { + is ValidationResult.Valid -> context.info( + stage = GraphExportStage.VALIDATION, + code = "minerva.graph.validation.passed", + message = "ComputeGraph validation passed before Minerva-specific checks." + ) + is ValidationResult.Invalid -> { + validation.errors.forEachIndexed { index, error -> + recordIssue( + issues = issues, + context = context, + kind = MinervaCompatibilityIssueKind.GRAPH_VALIDATION, + code = "minerva.compatibility.graph_invalid", + message = error, + remediation = "Fix the ComputeGraph structural validation error before exporting to Minerva.", + details = mapOf("errorIndex" to index.toString()) + ) + } + return report(options, context, issues, layerCount = 0, estimatedSramBytes = 0, estimatedFlashBytes = 0) + } + } + + val topology = try { + graph.getTopologicalOrder() + } catch (exception: Exception) { + recordIssue( + issues = issues, + context = context, + kind = MinervaCompatibilityIssueKind.UNSUPPORTED_TOPOLOGY, + code = "minerva.compatibility.topology_invalid", + message = exception.message ?: "Unable to determine graph topological order.", + remediation = "Use an acyclic sequential MLP graph." + ) + emptyList() + } + + validateOperations(topology, context, issues) + validateStaticShapes(graph, context, issues) + validateSequentialTopology(graph, context, issues) + validateActivationPlacement(graph, context, issues) + + val layerCount = topology.count { isLayerOperation(it.operationName) } + if (layerCount == 0) { + recordIssue( + issues = issues, + context = context, + kind = MinervaCompatibilityIssueKind.UNSUPPORTED_TOPOLOGY, + code = "minerva.compatibility.no_layers", + message = "Minerva phase-one export requires at least one dense, linear, or matmul layer.", + remediation = "Export a sequential MLP graph with at least one supported layer." + ) + } + + val estimatedSramBytes = estimateSramBytes(graph) + val estimatedFlashBytes = estimateFlashBytes(graph) + val maxSramBytes = (options.target.sramBytes - sramSafetyBytes).coerceAtLeast(0) + if (estimatedSramBytes > maxSramBytes) { + recordIssue( + issues = issues, + context = context, + kind = MinervaCompatibilityIssueKind.MEMORY_BUDGET_EXCEEDED, + code = "minerva.compatibility.sram_budget_exceeded", + message = "Estimated activation memory exceeds ${options.target.displayName} SRAM budget.", + remediation = "Reduce layer width, batch size, or target a Minerva configuration with more SRAM.", + details = mapOf( + "estimatedSramBytes" to estimatedSramBytes.toString(), + "targetSramBytes" to options.target.sramBytes.toString(), + "sramSafetyBytes" to sramSafetyBytes.toString() + ) + ) + } + + if (issues.isEmpty()) { + context.info( + stage = GraphExportStage.VALIDATION, + code = "minerva.compatibility.passed", + message = "ComputeGraph is compatible with the phase-one Minerva export contract.", + details = mapOf( + "layerCount" to layerCount.toString(), + "estimatedSramBytes" to estimatedSramBytes.toString(), + "estimatedFlashBytes" to estimatedFlashBytes.toString() + ) + ) + } + + return report(options, context, issues, layerCount, estimatedSramBytes, estimatedFlashBytes) + } + + private fun validateOperations( + nodes: List, + context: GraphExportContext, + issues: MutableList + ) { + nodes.forEach { node -> + if (!isSupportedOperation(node.operationName)) { + recordIssue( + issues = issues, + context = context, + kind = MinervaCompatibilityIssueKind.UNSUPPORTED_OPERATION, + code = "minerva.compatibility.unsupported_operation", + message = "Operation '${node.operationName}' is not supported by Minerva phase one.", + node = node, + remediation = "Use sequential MLP operations: matmul/dense/linear, add for bias, and relu/sigmoid/tanh activations.", + details = mapOf("operationType" to node.operationType) + ) + } + } + } + + private fun validateStaticShapes( + graph: ComputeGraph, + context: GraphExportContext, + issues: MutableList + ) { + graph.nodes.forEach { node -> + val specs = node.inputs.map { "input:${it.name}" to it } + node.outputs.map { "output:${it.name}" to it } + specs.forEach { (slot, spec) -> + val shape = spec.shape + when { + shape == null -> recordIssue( + issues = issues, + context = context, + kind = MinervaCompatibilityIssueKind.MISSING_SHAPE, + code = "minerva.compatibility.missing_shape", + message = "Tensor '$slot' on node '${node.id}' has a dynamic or missing shape.", + node = node, + remediation = "Provide fully static tensor shapes before exporting to Minerva.", + details = mapOf("tensor" to slot, "dtype" to spec.dtype) + ) + shape.isEmpty() || shape.any { it <= 0 } -> recordIssue( + issues = issues, + context = context, + kind = MinervaCompatibilityIssueKind.INVALID_SHAPE, + code = "minerva.compatibility.invalid_shape", + message = "Tensor '$slot' on node '${node.id}' has an invalid shape $shape.", + node = node, + remediation = "Use non-empty static shapes with positive dimensions.", + details = mapOf("tensor" to slot, "shape" to shape.joinToString("x")) + ) + } + } + } + graph.edges.forEach { edge -> + val shape = edge.tensorSpec.shape + if (shape == null) { + recordIssue( + issues = issues, + context = context, + kind = MinervaCompatibilityIssueKind.MISSING_SHAPE, + code = "minerva.compatibility.missing_edge_shape", + message = "Edge '${edge.id}' from '${edge.source.id}' to '${edge.destination.id}' has a dynamic or missing shape.", + node = edge.destination, + remediation = "Resolve all edge tensor shapes before exporting to Minerva.", + details = mapOf("edgeId" to edge.id) + ) + } else if (shape.isEmpty() || shape.any { it <= 0 }) { + recordIssue( + issues = issues, + context = context, + kind = MinervaCompatibilityIssueKind.INVALID_SHAPE, + code = "minerva.compatibility.invalid_edge_shape", + message = "Edge '${edge.id}' has an invalid shape $shape.", + node = edge.destination, + remediation = "Use non-empty static edge shapes with positive dimensions.", + details = mapOf("edgeId" to edge.id, "shape" to shape.joinToString("x")) + ) + } + } + } + + private fun validateSequentialTopology( + graph: ComputeGraph, + context: GraphExportContext, + issues: MutableList + ) { + val outputs = graph.getOutputNodes() + if (outputs.size != 1) { + recordIssue( + issues = issues, + context = context, + kind = MinervaCompatibilityIssueKind.UNSUPPORTED_TOPOLOGY, + code = "minerva.compatibility.output_count", + message = "Minerva phase one expects exactly one graph output, got ${outputs.size}.", + remediation = "Export one sequential MLP output tensor per Minerva model.", + details = mapOf("outputNodeIds" to outputs.joinToString(",") { it.id }) + ) + } + + graph.nodes.forEach { node -> + val incoming = graph.edges.count { it.destination == node } + val outgoing = graph.edges.count { it.source == node } + if (outgoing > 1) { + recordIssue( + issues = issues, + context = context, + kind = MinervaCompatibilityIssueKind.UNSUPPORTED_TOPOLOGY, + code = "minerva.compatibility.branching", + message = "Node '${node.id}' fans out to $outgoing consumers; Minerva phase one supports sequential MLP topology only.", + node = node, + remediation = "Remove branching or split the model into separate Minerva exports.", + details = mapOf("consumerCount" to outgoing.toString()) + ) + } + when (node.operationName.lowercase()) { + "matmul" -> requireIncoming(node, incoming, 2, context, issues) + "add" -> requireIncoming(node, incoming, 2, context, issues) + "relu", "sigmoid", "tanh" -> requireIncoming(node, incoming, 1, context, issues) + } + } + } + + private fun validateActivationPlacement( + graph: ComputeGraph, + context: GraphExportContext, + issues: MutableList + ) { + graph.nodes.filter { isActivationOperation(it.operationName) }.forEach { node -> + val producer = graph.getInputNodes(node).singleOrNull() + val producerName = producer?.operationName?.lowercase() + if (producerName !in activationProducerOperations) { + recordIssue( + issues = issues, + context = context, + kind = MinervaCompatibilityIssueKind.INCOMPATIBLE_ACTIVATION_PLACEMENT, + code = "minerva.compatibility.activation_placement", + message = "Activation node '${node.id}' must follow a dense, linear, matmul, or bias add layer.", + node = node, + remediation = "Place activations directly after a supported Minerva layer pattern.", + details = mapOf("producer" to (producerName ?: "none")) + ) + } + } + } + + private fun requireIncoming( + node: GraphNode, + actual: Int, + expected: Int, + context: GraphExportContext, + issues: MutableList + ) { + if (actual != expected) { + recordIssue( + issues = issues, + context = context, + kind = MinervaCompatibilityIssueKind.UNSUPPORTED_TOPOLOGY, + code = "minerva.compatibility.arity", + message = "Node '${node.id}' (${node.operationName}) expects $expected producer edge(s), got $actual.", + node = node, + remediation = "Use canonical sequential MLP fragments: matmul(input, weight), add(layer, bias), optional activation.", + details = mapOf("expected" to expected.toString(), "actual" to actual.toString()) + ) + } + } + + private fun estimateSramBytes(graph: ComputeGraph): Int { + val maxTensorBytes = (graph.nodes.flatMap { it.inputs + it.outputs } + graph.edges.map { it.tensorSpec }) + .mapNotNull { spec -> spec.shape?.let { tensorBytes(it) } } + .maxOrNull() ?: 0 + return maxTensorBytes * 2 + } + + private fun estimateFlashBytes(graph: ComputeGraph): Int { + return (graph.nodes.flatMap { it.inputs + it.outputs } + graph.edges.map { it.tensorSpec }) + .mapNotNull { spec -> spec.shape?.let { tensorBytes(it) } } + .sum() + } + + private fun tensorBytes(shape: List): Int { + return shape.fold(1) { acc, dim -> acc * dim } + } + + private fun report( + options: MinervaExportOptions, + context: GraphExportContext, + issues: List, + layerCount: Int, + estimatedSramBytes: Int, + estimatedFlashBytes: Int + ): MinervaCompatibilityReport { + return MinervaCompatibilityReport( + compatible = issues.isEmpty(), + diagnostics = context.diagnosticReport(), + issues = issues.toList(), + target = options.target, + quantization = options.quantization, + layerCount = layerCount, + estimatedSramBytes = estimatedSramBytes, + estimatedFlashBytes = estimatedFlashBytes, + metadata = context.metadata + ) + } + + private fun recordIssue( + issues: MutableList, + context: GraphExportContext, + kind: MinervaCompatibilityIssueKind, + code: String, + message: String, + remediation: String, + node: GraphNode? = null, + details: Map = emptyMap() + ) { + recordIssue( + issues = issues, + context = context, + kind = kind, + code = code, + message = message, + nodeId = node?.id, + operationName = node?.operationName, + remediation = remediation, + details = details + ) + } + + private fun recordIssue( + issues: MutableList, + context: GraphExportContext, + kind: MinervaCompatibilityIssueKind, + code: String, + message: String, + nodeId: String? = null, + operationName: String? = null, + remediation: String, + details: Map = emptyMap() + ) { + val issue = MinervaCompatibilityIssue( + kind = kind, + code = code, + message = message, + nodeId = nodeId, + operationName = operationName, + remediation = remediation, + details = details + ) + issues += issue + context.error( + stage = GraphExportStage.VALIDATION, + code = code, + message = message, + nodeId = nodeId, + operationName = operationName, + details = details + ("remediation" to remediation) + ) + } + + private fun isSupportedOperation(operationName: String): Boolean { + val name = operationName.lowercase() + return name == "input" || isLayerOperation(name) || name == "add" || isActivationOperation(name) + } + + private fun isLayerOperation(operationName: String): Boolean { + return operationName.lowercase() in layerOperations + } + + private fun isActivationOperation(operationName: String): Boolean { + return operationName.lowercase() in activationOperations + } + + public companion object { + public val layerOperations: Set = setOf("dense", "linear", "matmul") + public val activationOperations: Set = setOf("relu", "sigmoid", "tanh") + public val activationProducerOperations: Set = layerOperations + setOf("add") + } +} diff --git a/skainet-compile/skainet-compile-minerva/src/commonMain/kotlin/sk/ainet/compile/minerva/MinervaExportFacade.kt b/skainet-compile/skainet-compile-minerva/src/commonMain/kotlin/sk/ainet/compile/minerva/MinervaExportFacade.kt index ff7f37eea..04839b2d0 100644 --- a/skainet-compile/skainet-compile-minerva/src/commonMain/kotlin/sk/ainet/compile/minerva/MinervaExportFacade.kt +++ b/skainet-compile/skainet-compile-minerva/src/commonMain/kotlin/sk/ainet/compile/minerva/MinervaExportFacade.kt @@ -7,7 +7,6 @@ import sk.ainet.context.ExecutionContext import sk.ainet.lang.graph.ComputeGraph import sk.ainet.lang.graph.DefaultGraphExecutionContext import sk.ainet.lang.tape.toComputeGraph -import sk.ainet.lang.tensor.ops.ValidationResult import sk.ainet.lang.tensor.ops.VoidTensorOps import sk.ainet.tape.Execution @@ -20,7 +19,8 @@ import sk.ainet.tape.Execution * implementation issues. */ public class MinervaExportFacade @kotlin.jvm.JvmOverloads constructor( - public val backendName: String = MinervaExportBackend.backendName + public val backendName: String = MinervaExportBackend.backendName, + public val compatibilityValidator: MinervaCompatibilityValidator = MinervaCompatibilityValidator() ) { /** @@ -86,35 +86,19 @@ public class MinervaExportFacade @kotlin.jvm.JvmOverloads constructor( details = mapOf("nodes" to graph.nodes.size.toString()) ) - if (graph.nodes.isEmpty()) { - return graphValidationFailedResult( - options = options, - context = context, - errors = listOf("Minerva export requires at least one graph node.") - ) - } - - when (val validation = graph.validate()) { - is ValidationResult.Valid -> context.info( - stage = GraphExportStage.VALIDATION, - code = "minerva.graph.validation.passed", - message = "ComputeGraph validation passed before Minerva-specific checks." - ) - is ValidationResult.Invalid -> return graphValidationFailedResult( - options = options, - context = context, - errors = validation.errors - ) + val compatibilityReport = compatibilityValidator.validate(graph, options, context) + if (!compatibilityReport.compatible) { + return compatibilityValidationFailedResult(options, context, compatibilityReport) } val failure = MinervaExportFailure( kind = MinervaExportFailureKind.NOT_IMPLEMENTED, stage = GraphExportStage.LOWERING, code = "minerva.export.not_implemented", - message = "Minerva export API is scaffolded; compatibility validation, lowering, compiler invocation, packaging, and verification are implemented in follow-up issues.", + message = "Minerva export passed phase-one compatibility validation; lowering, compiler invocation, packaging, and verification are implemented in follow-up issues.", details = mapOf( - "nextStep" to "Implement MinervaCompatibilityValidator", - "issue" to "#691" + "nextStep" to "Implement MinervaGraphCanonicalizer", + "issue" to "#692" ) ) context.error( @@ -123,7 +107,7 @@ public class MinervaExportFacade @kotlin.jvm.JvmOverloads constructor( message = failure.message, details = failure.details ) - return failedResult(options, context, failure) + return failedResult(options, context, failure, compatibilityReport) } private fun unsupportedModelResult(model: Any, options: MinervaExportOptions): MinervaExportResult { @@ -165,31 +149,41 @@ public class MinervaExportFacade @kotlin.jvm.JvmOverloads constructor( return failedResult(options, context, failure) } - private fun graphValidationFailedResult( + private fun compatibilityValidationFailedResult( options: MinervaExportOptions, context: GraphExportContext, - errors: List + report: MinervaCompatibilityReport ): MinervaExportResult { + val firstIssue = report.issues.firstOrNull() + val details = mutableMapOf( + "issueCount" to report.issues.size.toString(), + "target" to report.target.compilerId, + "quantization" to report.quantization.compilerId + ) + if (firstIssue != null) { + details += mapOf( + "issueKind" to firstIssue.kind.name, + "remediation" to firstIssue.remediation + ) + firstIssue.nodeId?.let { details["nodeId"] = it } + firstIssue.operationName?.let { details["operationName"] = it } + details += firstIssue.details + } val failure = MinervaExportFailure( - kind = MinervaExportFailureKind.GRAPH_VALIDATION_FAILED, + kind = MinervaExportFailureKind.COMPATIBILITY_VALIDATION_FAILED, stage = GraphExportStage.VALIDATION, - code = "minerva.graph.validation_failed", - message = "ComputeGraph validation failed before Minerva-specific checks.", - details = errors.mapIndexed { index, error -> "error$index" to error }.toMap() - ) - context.error( - stage = failure.stage, - code = failure.code, - message = failure.message, - details = failure.details + code = firstIssue?.code ?: "minerva.compatibility.failed", + message = firstIssue?.message ?: "Minerva compatibility validation failed.", + details = details ) - return failedResult(options, context, failure) + return failedResult(options, context, failure, report) } private fun failedResult( options: MinervaExportOptions, context: GraphExportContext, - failure: MinervaExportFailure + failure: MinervaExportFailure, + compatibilityReport: MinervaCompatibilityReport? = null ): MinervaExportResult { return MinervaExportResult( options = options, @@ -197,7 +191,8 @@ public class MinervaExportFacade @kotlin.jvm.JvmOverloads constructor( diagnostics = context.diagnosticReport(), artifacts = context.artifacts, failure = failure, - metadata = context.metadata + metadata = context.metadata, + compatibilityReport = compatibilityReport ) } diff --git a/skainet-compile/skainet-compile-minerva/src/commonMain/kotlin/sk/ainet/compile/minerva/MinervaExportModels.kt b/skainet-compile/skainet-compile-minerva/src/commonMain/kotlin/sk/ainet/compile/minerva/MinervaExportModels.kt index b6bd88fcc..794fb9d6a 100644 --- a/skainet-compile/skainet-compile-minerva/src/commonMain/kotlin/sk/ainet/compile/minerva/MinervaExportModels.kt +++ b/skainet-compile/skainet-compile-minerva/src/commonMain/kotlin/sk/ainet/compile/minerva/MinervaExportModels.kt @@ -90,6 +90,7 @@ public enum class MinervaExportFailureKind { UNSUPPORTED_MODEL_TYPE, RECORDING_FAILED, GRAPH_VALIDATION_FAILED, + COMPATIBILITY_VALIDATION_FAILED, NOT_IMPLEMENTED } @@ -130,6 +131,65 @@ public data class MinervaExportBundle( } } +/** + * Stable categories for Minerva compatibility findings. + */ +public enum class MinervaCompatibilityIssueKind { + GRAPH_VALIDATION, + UNSUPPORTED_OPERATION, + UNSUPPORTED_TOPOLOGY, + MISSING_SHAPE, + INVALID_SHAPE, + INCOMPATIBLE_ACTIVATION_PLACEMENT, + MEMORY_BUDGET_EXCEEDED, + UNSUPPORTED_QUANTIZATION +} + +/** + * Backend-specific compatibility issue that also appears as a graph-export diagnostic. + */ +public data class MinervaCompatibilityIssue( + public val kind: MinervaCompatibilityIssueKind, + public val code: String, + public val message: String, + public val nodeId: String? = null, + public val operationName: String? = null, + public val remediation: String, + public val details: Map = emptyMap() +) { + init { + require(code.isNotBlank()) { "compatibility issue code cannot be blank" } + require(message.isNotBlank()) { "compatibility issue message cannot be blank" } + require(remediation.isNotBlank()) { "compatibility issue remediation cannot be blank" } + } +} + +/** + * Phase-one Minerva compatibility report. + */ +public data class MinervaCompatibilityReport( + public val compatible: Boolean, + public val diagnostics: GraphExportDiagnosticReport, + public val issues: List, + public val target: MinervaTarget, + public val quantization: MinervaQuantization, + public val layerCount: Int, + public val estimatedSramBytes: Int, + public val estimatedFlashBytes: Int, + public val metadata: Map = emptyMap() +) { + public val failed: Boolean + get() = !compatible + + public fun requireCompatible(): MinervaCompatibilityReport { + if (!compatible) { + val summary = issues.joinToString("; ") { "${it.code}: ${it.message}" } + error("Minerva compatibility validation failed: $summary") + } + return this + } +} + /** * Public result shape for Minerva export attempts. */ @@ -140,7 +200,8 @@ public data class MinervaExportResult( public val diagnostics: GraphExportDiagnosticReport = GraphExportDiagnosticReport.empty(), public val artifacts: List = emptyList(), public val failure: MinervaExportFailure? = null, - public val metadata: Map = emptyMap() + public val metadata: Map = emptyMap(), + public val compatibilityReport: MinervaCompatibilityReport? = null ) { init { require(status != GraphExportStatus.SUCCESS || bundle != null) { diff --git a/skainet-compile/skainet-compile-minerva/src/commonTest/kotlin/sk/ainet/compile/minerva/MinervaCompatibilityValidatorTest.kt b/skainet-compile/skainet-compile-minerva/src/commonTest/kotlin/sk/ainet/compile/minerva/MinervaCompatibilityValidatorTest.kt new file mode 100644 index 000000000..fd2c2c599 --- /dev/null +++ b/skainet-compile/skainet-compile-minerva/src/commonTest/kotlin/sk/ainet/compile/minerva/MinervaCompatibilityValidatorTest.kt @@ -0,0 +1,121 @@ +package sk.ainet.compile.minerva + +import kotlin.test.Test +import kotlin.test.assertEquals +import kotlin.test.assertFalse +import kotlin.test.assertTrue +import sk.ainet.lang.graph.DefaultComputeGraph + +class MinervaCompatibilityValidatorTest { + + @Test + fun supportedMlpGraphPassesCompatibilityValidation() { + val report = MinervaCompatibilityValidator().validate( + graph = validMinervaMlpGraph(), + options = minervaTestOptions() + ) + + assertTrue(report.compatible) + assertFalse(report.failed) + assertEquals(1, report.layerCount) + assertTrue(report.issues.isEmpty()) + assertTrue(report.estimatedSramBytes > 0) + assertTrue(report.diagnostics.infos.any { it.code == "minerva.compatibility.passed" }) + } + + @Test + fun emptyGraphFailsBeforeCompilerInvocation() { + val report = MinervaCompatibilityValidator().validate( + graph = DefaultComputeGraph(), + options = minervaTestOptions() + ) + + assertFalse(report.compatible) + assertTrue(report.diagnostics.hasErrors) + assertTrue( + report.issues.any { + it.kind == MinervaCompatibilityIssueKind.GRAPH_VALIDATION && + it.code == "minerva.compatibility.empty_graph" + } + ) + } + + @Test + fun unsupportedOperationNamesNodeAndRemediation() { + val report = MinervaCompatibilityValidator().validate( + graph = unsupportedMinervaOperationGraph(), + options = minervaTestOptions() + ) + val issue = report.issues.first { + it.kind == MinervaCompatibilityIssueKind.UNSUPPORTED_OPERATION + } + + assertFalse(report.compatible) + assertEquals("conv", issue.nodeId) + assertEquals("conv1d", issue.operationName) + assertTrue(issue.remediation.contains("sequential MLP")) + } + + @Test + fun unsupportedTopologyNamesBranchingNode() { + val report = MinervaCompatibilityValidator().validate( + graph = branchingMinervaGraph(), + options = minervaTestOptions() + ) + val issue = report.issues.first { + it.kind == MinervaCompatibilityIssueKind.UNSUPPORTED_TOPOLOGY && + it.code == "minerva.compatibility.branching" + } + + assertFalse(report.compatible) + assertEquals("input", issue.nodeId) + assertEquals("2", issue.details["consumerCount"]) + } + + @Test + fun missingStaticShapesFailCompatibilityValidation() { + val report = MinervaCompatibilityValidator().validate( + graph = missingShapeMinervaGraph(), + options = minervaTestOptions() + ) + + assertFalse(report.compatible) + assertTrue( + report.issues.any { + it.kind == MinervaCompatibilityIssueKind.MISSING_SHAPE && + it.nodeId == "input" + } + ) + assertTrue(report.diagnostics.errors.any { it.code == "minerva.compatibility.missing_shape" }) + } + + @Test + fun activationPlacementMustFollowSupportedLayerPattern() { + val report = MinervaCompatibilityValidator().validate( + graph = activationBeforeLayerGraph(), + options = minervaTestOptions() + ) + val issue = report.issues.first { + it.kind == MinervaCompatibilityIssueKind.INCOMPATIBLE_ACTIVATION_PLACEMENT + } + + assertFalse(report.compatible) + assertEquals("relu", issue.nodeId) + assertEquals("input", issue.details["producer"]) + } + + @Test + fun targetMemoryOverflowFailsCompatibilityValidation() { + val report = MinervaCompatibilityValidator().validate( + graph = validMinervaMlpGraph(inputWidth = 2048, outputWidth = 16), + options = minervaTestOptions() + ) + val issue = report.issues.first { + it.kind == MinervaCompatibilityIssueKind.MEMORY_BUDGET_EXCEEDED + } + + assertFalse(report.compatible) + assertEquals("2048", issue.details["targetSramBytes"]) + assertTrue(report.estimatedSramBytes > MinervaTarget.ATMEGA328P.sramBytes) + } +} diff --git a/skainet-compile/skainet-compile-minerva/src/commonTest/kotlin/sk/ainet/compile/minerva/MinervaExportFacadeTest.kt b/skainet-compile/skainet-compile-minerva/src/commonTest/kotlin/sk/ainet/compile/minerva/MinervaExportFacadeTest.kt index 3d3e1c8f5..0f04a9639 100644 --- a/skainet-compile/skainet-compile-minerva/src/commonTest/kotlin/sk/ainet/compile/minerva/MinervaExportFacadeTest.kt +++ b/skainet-compile/skainet-compile-minerva/src/commonTest/kotlin/sk/ainet/compile/minerva/MinervaExportFacadeTest.kt @@ -8,17 +8,13 @@ import kotlin.test.assertNotNull import kotlin.test.assertTrue import sk.ainet.compile.export.GraphExportStatus import sk.ainet.lang.graph.DefaultComputeGraph -import sk.ainet.lang.graph.GraphNode -import sk.ainet.lang.tensor.ops.InputOperation -import sk.ainet.lang.tensor.ops.TensorSpec -import sk.ainet.lang.types.DType class MinervaExportFacadeTest { @Test fun createsFacadeAndDefaultOptions() { val facade = MinervaExportFacade() - val options = testOptions() + val options = minervaTestOptions() assertEquals(MinervaExportBackend.backendName, facade.backendName) assertEquals(MinervaTarget.ATMEGA328P, options.target) @@ -29,35 +25,37 @@ class MinervaExportFacadeTest { @Test fun rejectsInvalidOptionsWithClearMessages() { val outputError = assertFailsWith { - testOptions(outputDir = "") + minervaTestOptions(outputDir = "") } assertTrue(outputError.message?.contains("outputDir cannot be blank") == true) val projectError = assertFailsWith { - testOptions(projectName = "nested/project") + minervaTestOptions(projectName = "nested/project") } assertTrue(projectError.message?.contains("simple project directory name") == true) } @Test fun exportGraphRejectsEmptyGraphBeforePlaceholderStage() { - val result = MinervaExportFacade().exportGraph(DefaultComputeGraph(), testOptions()) + val result = MinervaExportFacade().exportGraph(DefaultComputeGraph(), minervaTestOptions()) assertEquals(GraphExportStatus.FAILED, result.status) assertFalse(result.succeeded) - assertEquals(MinervaExportFailureKind.GRAPH_VALIDATION_FAILED, result.failure?.kind) + assertEquals(MinervaExportFailureKind.COMPATIBILITY_VALIDATION_FAILED, result.failure?.kind) + assertEquals(MinervaCompatibilityIssueKind.GRAPH_VALIDATION, result.compatibilityReport?.issues?.first()?.kind) assertTrue(result.diagnostics.hasErrors) - assertTrue(result.failure?.details?.values?.any { it.contains("at least one graph node") } == true) + assertTrue(result.failure?.message?.contains("at least one graph node") == true) } @Test fun exportGraphReturnsNotImplementedForValidatedGraph() { - val result = MinervaExportFacade().exportGraph(singleInputGraph(), testOptions()) + val result = MinervaExportFacade().exportGraph(validMinervaMlpGraph(), minervaTestOptions()) assertEquals(GraphExportStatus.FAILED, result.status) assertEquals(MinervaExportFailureKind.NOT_IMPLEMENTED, result.failure?.kind) assertEquals("minerva.export.not_implemented", result.failure?.code) assertTrue(result.diagnostics.infos.any { it.code == "minerva.graph.validation.passed" }) + assertTrue(result.compatibilityReport?.compatible == true) assertTrue(result.metadata["target"] == MinervaTarget.ATMEGA328P.compilerId) assertFailsWith { result.requireSuccess() @@ -66,15 +64,16 @@ class MinervaExportFacadeTest { @Test fun exportModelAcceptsComputeGraphFastPath() { - val graph = singleInputGraph() - val result = MinervaExportFacade().exportModel(graph, testOptions()) + val graph = validMinervaMlpGraph() + val result = MinervaExportFacade().exportModel(graph, minervaTestOptions()) assertEquals(MinervaExportFailureKind.NOT_IMPLEMENTED, result.failure?.kind) + assertTrue(result.compatibilityReport?.compatible == true) } @Test fun exportModelReportsUnsupportedModelWithoutForwardPass() { - val result = MinervaExportFacade().exportModel("not-a-graph", testOptions()) + val result = MinervaExportFacade().exportModel("not-a-graph", minervaTestOptions()) assertEquals(GraphExportStatus.FAILED, result.status) assertEquals(MinervaExportFailureKind.UNSUPPORTED_MODEL_TYPE, result.failure?.kind) @@ -86,38 +85,35 @@ class MinervaExportFacadeTest { val result = MinervaExportFacade().exportModel( model = object {}, forwardPass = { }, - options = testOptions(projectName = "RecordedModel") + options = minervaTestOptions(projectName = "RecordedModel") ) assertEquals(GraphExportStatus.FAILED, result.status) val failure = assertNotNull(result.failure) assertTrue( - failure.kind == MinervaExportFailureKind.GRAPH_VALIDATION_FAILED || + failure.kind == MinervaExportFailureKind.COMPATIBILITY_VALIDATION_FAILED || failure.kind == MinervaExportFailureKind.RECORDING_FAILED ) } - private fun testOptions( - outputDir: String = "build/minerva", - projectName: String = "TinyMlp" - ): MinervaExportOptions { - return MinervaExportOptions( - outputDir = outputDir, - projectName = projectName, - metadata = mapOf("test" to "true") + @Test + fun exportGraphIncludesCompatibilityReportForUnsupportedGraph() { + val result = MinervaExportFacade().exportGraph( + graph = unsupportedMinervaOperationGraph(), + options = minervaTestOptions() ) - } - private fun singleInputGraph(): DefaultComputeGraph { - val graph = DefaultComputeGraph() - graph.addNode( - GraphNode( - id = "input", - operation = InputOperation(), - inputs = emptyList(), - outputs = listOf(TensorSpec("x", listOf(1, 4), "FP32")) - ) + assertEquals(GraphExportStatus.FAILED, result.status) + assertEquals(MinervaExportFailureKind.COMPATIBILITY_VALIDATION_FAILED, result.failure?.kind) + val report = assertNotNull(result.compatibilityReport) + assertFalse(report.compatible) + assertTrue( + report.issues.any { + it.kind == MinervaCompatibilityIssueKind.UNSUPPORTED_OPERATION && + it.nodeId == "conv" && + it.operationName == "conv1d" + } ) - return graph + assertEquals("conv", result.failure?.details?.get("nodeId")) } } diff --git a/skainet-compile/skainet-compile-minerva/src/commonTest/kotlin/sk/ainet/compile/minerva/MinervaGraphFixtures.kt b/skainet-compile/skainet-compile-minerva/src/commonTest/kotlin/sk/ainet/compile/minerva/MinervaGraphFixtures.kt new file mode 100644 index 000000000..82e7eb2c1 --- /dev/null +++ b/skainet-compile/skainet-compile-minerva/src/commonTest/kotlin/sk/ainet/compile/minerva/MinervaGraphFixtures.kt @@ -0,0 +1,187 @@ +package sk.ainet.compile.minerva + +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.GenericOperation +import sk.ainet.lang.tensor.ops.InputOperation +import sk.ainet.lang.tensor.ops.MatmulOperation +import sk.ainet.lang.tensor.ops.ReluOperation +import sk.ainet.lang.tensor.ops.TensorSpec +import sk.ainet.lang.types.DType + +internal fun minervaTestOptions( + outputDir: String = "build/minerva", + projectName: String = "TinyMlp" +): MinervaExportOptions { + return MinervaExportOptions( + outputDir = outputDir, + projectName = projectName, + metadata = mapOf("test" to "true") + ) +} + +internal fun validMinervaMlpGraph( + inputWidth: Int = 4, + outputWidth: Int = 3 +): DefaultComputeGraph { + val xSpec = spec("x", 1, inputWidth) + val wSpec = spec("w", inputWidth, outputWidth) + val matmulSpec = spec("matmul", 1, outputWidth) + val biasSpec = spec("bias", 1, outputWidth) + val addSpec = spec("biased", 1, outputWidth) + val ySpec = spec("y", 1, outputWidth) + + val x = inputNode("input", xSpec) + val w = inputNode("weight", wSpec) + val matmul = GraphNode( + id = "matmul", + operation = MatmulOperation(), + inputs = listOf(xSpec, wSpec), + outputs = listOf(matmulSpec) + ) + val bias = inputNode("bias", biasSpec) + val add = GraphNode( + id = "bias_add", + operation = AddOperation(), + inputs = listOf(matmulSpec, biasSpec), + outputs = listOf(addSpec) + ) + val relu = GraphNode( + id = "relu", + operation = ReluOperation(), + inputs = listOf(addSpec), + outputs = listOf(ySpec) + ) + + return graphOf( + nodes = listOf(x, w, matmul, bias, add, relu), + edges = listOf( + edge("x_to_matmul", x, matmul, xSpec, destinationInputIndex = 0), + edge("w_to_matmul", w, matmul, wSpec, destinationInputIndex = 1), + edge("matmul_to_add", matmul, add, matmulSpec, destinationInputIndex = 0), + edge("bias_to_add", bias, add, biasSpec, destinationInputIndex = 1), + edge("add_to_relu", add, relu, addSpec) + ) + ) +} + +internal fun unsupportedMinervaOperationGraph(): DefaultComputeGraph { + val inputSpec = spec("x", 1, 4) + val outputSpec = spec("conv_out", 1, 4) + val input = inputNode("input", inputSpec) + val conv = GraphNode( + id = "conv", + operation = GenericOperation("conv1d", type = "nn"), + inputs = listOf(inputSpec), + outputs = listOf(outputSpec) + ) + + return graphOf( + nodes = listOf(input, conv), + edges = listOf(edge("input_to_conv", input, conv, inputSpec)) + ) +} + +internal fun branchingMinervaGraph(): DefaultComputeGraph { + val inputSpec = spec("x", 1, 4) + val reluSpec = spec("relu", 1, 4) + val branchSpec = spec("branch", 1, 4) + val input = inputNode("input", inputSpec) + val relu = GraphNode( + id = "relu_a", + operation = ReluOperation(), + inputs = listOf(inputSpec), + outputs = listOf(reluSpec) + ) + val sigmoid = GraphNode( + id = "relu_b", + operation = ReluOperation(), + inputs = listOf(inputSpec), + outputs = listOf(branchSpec) + ) + + return graphOf( + nodes = listOf(input, relu, sigmoid), + edges = listOf( + edge("input_to_relu_a", input, relu, inputSpec), + edge("input_to_relu_b", input, sigmoid, inputSpec) + ) + ) +} + +internal fun missingShapeMinervaGraph(): DefaultComputeGraph { + val xSpec = TensorSpec("x", null, "Float32") + val wSpec = spec("w", 4, 3) + val outputSpec = spec("matmul", 1, 3) + val input = inputNode("input", xSpec) + val weight = inputNode("weight", wSpec) + val matmul = GraphNode( + id = "matmul", + operation = MatmulOperation(), + inputs = listOf(xSpec, wSpec), + outputs = listOf(outputSpec) + ) + + return graphOf( + nodes = listOf(input, weight, matmul), + edges = listOf( + edge("x_to_matmul", input, matmul, xSpec, destinationInputIndex = 0), + edge("w_to_matmul", weight, matmul, wSpec, destinationInputIndex = 1) + ) + ) +} + +internal fun activationBeforeLayerGraph(): DefaultComputeGraph { + val inputSpec = spec("x", 1, 4) + val outputSpec = spec("relu", 1, 4) + val input = inputNode("input", inputSpec) + val relu = GraphNode( + id = "relu", + operation = ReluOperation(), + inputs = listOf(inputSpec), + outputs = listOf(outputSpec) + ) + + return graphOf( + nodes = listOf(input, relu), + edges = listOf(edge("input_to_relu", input, relu, inputSpec)) + ) +} + +private fun inputNode(id: String, output: TensorSpec): GraphNode { + return GraphNode( + id = id, + operation = InputOperation(), + inputs = emptyList(), + outputs = listOf(output) + ) +} + +private fun spec(name: String, vararg shape: Int): TensorSpec { + return TensorSpec(name, shape.toList(), "Float32") +} + +private fun graphOf(nodes: List, edges: List): DefaultComputeGraph { + val graph = DefaultComputeGraph() + nodes.forEach { graph.addNode(it) } + edges.forEach { graph.addEdge(it) } + return graph +} + +private fun edge( + id: String, + source: GraphNode, + destination: GraphNode, + spec: TensorSpec, + destinationInputIndex: Int = 0 +): GraphEdge { + return GraphEdge( + id = id, + source = source, + destination = destination, + destinationInputIndex = destinationInputIndex, + tensorSpec = spec + ) +}