Skip to content

Commit 84e6121

Browse files
Merge pull request #269 from SKaiNET-developers/feature/267-arduino-plain-c
Add support for arduino c code generating
2 parents 98096df + f973765 commit 84e6121

65 files changed

Lines changed: 8052 additions & 942 deletions

File tree

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

README.md

Lines changed: 14 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -72,12 +72,25 @@ println(ds.describe())
7272
### SKaiNET is Compiler
7373

7474
- MLIR/StableHLO based lowering (modules provided in `SKaiNET-compile-*`)
75-
7675
```kotlin
7776
// Illustrative: export graph to JSON/StableHLO IR
7877
val ir = Compile.toStableHlo(model)
7978
println(ir.pretty())
8079
```
80+
- **Arduino C Code Generation**: Export models to standalone, optimized C99 code with static memory allocation.
81+
82+
```kotlin
83+
// Export model to an Arduino library
84+
val facade = CCodegenFacade()
85+
facade.exportToArduinoLibrary(
86+
model = model,
87+
forwardPass = { ctx -> model.forward(input, ctx) },
88+
outputPath = "build/arduino",
89+
libraryName = "MyModel"
90+
)
91+
```
92+
93+
Read the [Deep Technical Explanation](docs/arduino-c-codegen.md) for more details.
8194

8295
### SKaiNET is for Developers
8396

docs/SKaiNET-compiler.svg

Lines changed: 76 additions & 24 deletions
Loading

docs/arduino-c-codegen.md

Lines changed: 75 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,75 @@
1+
# Arduino C Code Generation
2+
3+
SKaiNET provides a specialized compiler backend for exporting trained neural networks to highly optimized, standalone C99 code suitable for microcontrollers like Arduino.
4+
5+
## Overview
6+
7+
The Arduino C code generation process transforms a high-level Kotlin model into a memory-efficient C implementation. It prioritizes static memory allocation, minimal overhead, and numerical consistency with the original model.
8+
9+
### Codegen Pipeline
10+
11+
```mermaid
12+
graph TD
13+
A[Kotlin Model] --> B[Recording Pass]
14+
B --> C[Execution Tape]
15+
C --> D[Compute Graph]
16+
D --> E[Graph Validation]
17+
E --> F[Memory Layout Calculation]
18+
F --> G[C Code Emission]
19+
G --> H[Arduino Library Packaging]
20+
H --> I[Generated .h/.c files]
21+
```
22+
23+
## Technical Deep Dive
24+
25+
### 1. Tape-based Tracing
26+
Instead of static analysis of the Kotlin code, SKaiNET uses a dynamic tracing mechanism. When you call `exportToArduinoLibrary`, the framework executes a single forward pass of your model using a specialized `RecordingContext`.
27+
- Every operation (Dense, ReLU, etc.) is recorded onto an **Execution Tape**.
28+
- This approach handles Kotlin's language features (loops, conditionals) naturally, as it only records the actual operations that were executed.
29+
30+
### 2. Compute Graph Construction
31+
The execution tape is converted into a directed acyclic graph (DAG) called `ComputeGraph`.
32+
- Nodes represent operations (Ops).
33+
- Edges represent data flow (Tensors).
34+
- During this phase, the compiler performs **Shape Inference** to ensure every tensor has a fixed, known size.
35+
36+
### 3. Static Memory Management
37+
Microcontrollers typically have very limited RAM and lack robust heap management. SKaiNET uses a **Ping-Pong Buffer Strategy** to eliminate dynamic memory allocation (`malloc`/`free`) during inference.
38+
39+
#### Ping-Pong Buffer Strategy
40+
The compiler calculates the maximum size required for any intermediate tensor in the graph and allocates exactly two static buffers of that size.
41+
42+
```mermaid
43+
sequenceDiagram
44+
participant I as Input
45+
participant B1 as Buffer A
46+
participant B2 as Buffer B
47+
participant O as Output
48+
49+
I->>B1: Layer 1 (Input -> A)
50+
B1->>B2: Layer 2 (A -> B)
51+
B2->>B1: Layer 3 (B -> A)
52+
B1->>O: Layer 4 (A -> Output)
53+
```
54+
55+
- **Buffer Reuse**: Instead of allocating space for every layer's output, buffers are reused.
56+
- **Direct Output Optimization**: The first layer reads from the input pointer, and the last layer writes directly to the output pointer, avoiding unnecessary copies.
57+
58+
### 4. Code Generation (Emission)
59+
The `CCodeGenerator` emits C99-compatible code using templates.
60+
- **Weights & Biases**: Extracted from the trained Kotlin model and serialized as `static const float` arrays. This places them in Flash memory (PROGMEM) on many microcontrollers, saving precious RAM.
61+
- **Kernel Implementation**: Operations like `Dense` (Linear) are implemented as optimized nested loops.
62+
- **Header Generation**: Produces a clean API for the user:
63+
```c
64+
int model_inference(const float* input, float* output);
65+
```
66+
67+
### 5. Validation
68+
The generator performs post-generation validation:
69+
- **Static Allocation Check**: Ensures no dynamic allocation is present in the generated source.
70+
- **Buffer Alternation Check**: Verifies that the ping-pong strategy is correctly implemented without data races or overwrites.
71+
72+
## Performance and Constraints
73+
- **Floating Point**: Currently optimized for `FP32`.
74+
- **Supported Ops**: `Dense`, `ReLU`, `Sigmoid`, `Tanh`, `Add`, `MatMul`.
75+
- **Memory**: Total memory consumption is `TotalWeights + 2 * MaxIntermediateTensor`.

settings.gradle.kts

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -32,6 +32,7 @@ include("skainet-lang:skainet-kan")
3232
include("skainet-compile:skainet-compile-core")
3333
include("skainet-compile:skainet-compile-dag")
3434
include("skainet-compile:skainet-compile-json")
35+
include("skainet-compile:skainet-compile-c")
3536

3637
// ====== BACKENDS
3738
include("skainet-backends:skainet-backend-cpu")
Lines changed: 133 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,133 @@
1+
public final class sk/ainet/compile/c/ArduinoCodegenIntegration {
2+
public fun <init> ()V
3+
public final fun generateArduinoLibrary (Lsk/ainet/lang/graph/ComputeGraph;Ljava/lang/String;Ljava/lang/String;)Lsk/ainet/compile/c/ArduinoLibraryResult;
4+
}
5+
6+
public final class sk/ainet/compile/c/ArduinoLibraryPackager {
7+
public fun <init> ()V
8+
public final fun createLibraryStructure (Ljava/lang/String;Ljava/lang/String;Ljava/lang/String;Ljava/lang/String;Lsk/ainet/compile/c/MemoryLayout;[I[I)Lsk/ainet/compile/c/ArduinoLibraryResult;
9+
}
10+
11+
public final class sk/ainet/compile/c/ArduinoLibraryResult {
12+
public fun <init> (Ljava/lang/String;Lsk/ainet/compile/c/MemoryLayout;Ljava/util/List;Ljava/util/List;)V
13+
public final fun component1 ()Ljava/lang/String;
14+
public final fun component2 ()Lsk/ainet/compile/c/MemoryLayout;
15+
public final fun component3 ()Ljava/util/List;
16+
public final fun component4 ()Ljava/util/List;
17+
public final fun copy (Ljava/lang/String;Lsk/ainet/compile/c/MemoryLayout;Ljava/util/List;Ljava/util/List;)Lsk/ainet/compile/c/ArduinoLibraryResult;
18+
public static synthetic fun copy$default (Lsk/ainet/compile/c/ArduinoLibraryResult;Ljava/lang/String;Lsk/ainet/compile/c/MemoryLayout;Ljava/util/List;Ljava/util/List;ILjava/lang/Object;)Lsk/ainet/compile/c/ArduinoLibraryResult;
19+
public fun equals (Ljava/lang/Object;)Z
20+
public final fun getGeneratedFiles ()Ljava/util/List;
21+
public final fun getLibraryPath ()Ljava/lang/String;
22+
public final fun getMemoryRequirements ()Lsk/ainet/compile/c/MemoryLayout;
23+
public final fun getSupportedOperations ()Ljava/util/List;
24+
public fun hashCode ()I
25+
public fun toString ()Ljava/lang/String;
26+
}
27+
28+
public final class sk/ainet/compile/c/CCodeGenerator {
29+
public static final field Companion Lsk/ainet/compile/c/CCodeGenerator$Companion;
30+
public fun <init> (Lsk/ainet/lang/graph/ComputeGraph;)V
31+
public final fun calculateMemoryRequirements ()Lsk/ainet/compile/c/MemoryLayout;
32+
public final fun extractWeights ()Ljava/util/List;
33+
public final fun generateActivationFunction (Lsk/ainet/lang/graph/GraphNode;)Lsk/ainet/compile/c/LayerCode;
34+
public final fun generateActivationFunctionWithAccuracy (Lsk/ainet/lang/graph/GraphNode;)Lsk/ainet/compile/c/LayerCode;
35+
public final fun generateAllLayers ()Ljava/util/List;
36+
public final fun generateDenseLayer (Lsk/ainet/lang/graph/GraphNode;)Lsk/ainet/compile/c/LayerCode;
37+
public final fun generateDenseLayerWithAccuracy (Lsk/ainet/lang/graph/GraphNode;)Lsk/ainet/compile/c/LayerCode;
38+
public final fun generateTransposeLayer (Lsk/ainet/lang/graph/GraphNode;)Lsk/ainet/compile/c/LayerCode;
39+
public final fun validateGeneratedCodeBufferAlternation (Ljava/lang/String;)Lsk/ainet/lang/tensor/ops/ValidationResult;
40+
public final fun validateGeneratedCodeStaticAllocation (Ljava/lang/String;)Lsk/ainet/lang/tensor/ops/ValidationResult;
41+
public final fun validateGraph ()Lsk/ainet/lang/tensor/ops/ValidationResult;
42+
public final fun validateMemoryManagement ()Lsk/ainet/lang/tensor/ops/ValidationResult;
43+
public final fun validateOperations ()Lsk/ainet/lang/tensor/ops/ValidationResult;
44+
}
45+
46+
public final class sk/ainet/compile/c/CCodeGenerator$Companion {
47+
}
48+
49+
public final class sk/ainet/compile/c/CCodegenFacade {
50+
public fun <init> ()V
51+
public final fun exportGraphToArduinoLibrary (Lsk/ainet/lang/graph/ComputeGraph;Ljava/lang/String;Ljava/lang/String;)Lsk/ainet/compile/c/ArduinoLibraryResult;
52+
public final fun exportToArduinoLibrary (Ljava/lang/Object;Ljava/lang/String;Ljava/lang/String;)Lsk/ainet/compile/c/ArduinoLibraryResult;
53+
public final fun exportToArduinoLibrary (Ljava/lang/Object;Lkotlin/jvm/functions/Function1;Ljava/lang/String;Ljava/lang/String;)Lsk/ainet/compile/c/ArduinoLibraryResult;
54+
public static synthetic fun exportToArduinoLibrary$default (Lsk/ainet/compile/c/CCodegenFacade;Ljava/lang/Object;Ljava/lang/String;Ljava/lang/String;ILjava/lang/Object;)Lsk/ainet/compile/c/ArduinoLibraryResult;
55+
public static synthetic fun exportToArduinoLibrary$default (Lsk/ainet/compile/c/CCodegenFacade;Ljava/lang/Object;Lkotlin/jvm/functions/Function1;Ljava/lang/String;Ljava/lang/String;ILjava/lang/Object;)Lsk/ainet/compile/c/ArduinoLibraryResult;
56+
}
57+
58+
public final class sk/ainet/compile/c/LayerCode {
59+
public fun <init> (Ljava/lang/String;Ljava/lang/String;[I[ILjava/lang/String;)V
60+
public final fun component1 ()Ljava/lang/String;
61+
public final fun component2 ()Ljava/lang/String;
62+
public final fun component3 ()[I
63+
public final fun component4 ()[I
64+
public final fun component5 ()Ljava/lang/String;
65+
public final fun copy (Ljava/lang/String;Ljava/lang/String;[I[ILjava/lang/String;)Lsk/ainet/compile/c/LayerCode;
66+
public static synthetic fun copy$default (Lsk/ainet/compile/c/LayerCode;Ljava/lang/String;Ljava/lang/String;[I[ILjava/lang/String;ILjava/lang/Object;)Lsk/ainet/compile/c/LayerCode;
67+
public fun equals (Ljava/lang/Object;)Z
68+
public final fun getCodeFragment ()Ljava/lang/String;
69+
public final fun getInputShape ()[I
70+
public final fun getLayerName ()Ljava/lang/String;
71+
public final fun getOperationType ()Ljava/lang/String;
72+
public final fun getOutputShape ()[I
73+
public fun hashCode ()I
74+
public fun toString ()Ljava/lang/String;
75+
}
76+
77+
public final class sk/ainet/compile/c/MemoryLayout {
78+
public fun <init> (IIILjava/util/List;)V
79+
public final fun component1 ()I
80+
public final fun component2 ()I
81+
public final fun component3 ()I
82+
public final fun component4 ()Ljava/util/List;
83+
public final fun copy (IIILjava/util/List;)Lsk/ainet/compile/c/MemoryLayout;
84+
public static synthetic fun copy$default (Lsk/ainet/compile/c/MemoryLayout;IIILjava/util/List;ILjava/lang/Object;)Lsk/ainet/compile/c/MemoryLayout;
85+
public fun equals (Ljava/lang/Object;)Z
86+
public final fun getBufferSizes ()Ljava/util/List;
87+
public final fun getMaxIntermediateSize ()I
88+
public final fun getTotalMemoryRequired ()I
89+
public final fun getTotalWeightSize ()I
90+
public fun hashCode ()I
91+
public fun toString ()Ljava/lang/String;
92+
}
93+
94+
public final class sk/ainet/compile/c/Platform_androidKt {
95+
public static final fun platformCreateDirectory (Ljava/lang/String;)V
96+
public static final fun platformWriteFile (Ljava/lang/String;Ljava/lang/String;)V
97+
}
98+
99+
public final class sk/ainet/compile/c/WeightArray {
100+
public fun <init> (Ljava/lang/String;[F[IZ)V
101+
public synthetic fun <init> (Ljava/lang/String;[F[IZILkotlin/jvm/internal/DefaultConstructorMarker;)V
102+
public final fun component1 ()Ljava/lang/String;
103+
public final fun component2 ()[F
104+
public final fun component3 ()[I
105+
public final fun component4 ()Z
106+
public final fun copy (Ljava/lang/String;[F[IZ)Lsk/ainet/compile/c/WeightArray;
107+
public static synthetic fun copy$default (Lsk/ainet/compile/c/WeightArray;Ljava/lang/String;[F[IZILjava/lang/Object;)Lsk/ainet/compile/c/WeightArray;
108+
public fun equals (Ljava/lang/Object;)Z
109+
public final fun getName ()Ljava/lang/String;
110+
public final fun getShape ()[I
111+
public final fun getValues ()[F
112+
public fun hashCode ()I
113+
public final fun isWeight ()Z
114+
public fun toString ()Ljava/lang/String;
115+
}
116+
117+
public final class sk/ainet/compile/c/templates/ExampleTemplate {
118+
public static final field INSTANCE Lsk/ainet/compile/c/templates/ExampleTemplate;
119+
public final fun generate (Ljava/lang/String;[I[ILjava/lang/String;)Ljava/lang/String;
120+
public static synthetic fun generate$default (Lsk/ainet/compile/c/templates/ExampleTemplate;Ljava/lang/String;[I[ILjava/lang/String;ILjava/lang/Object;)Ljava/lang/String;
121+
public final fun generateMinimal (Ljava/lang/String;II)Ljava/lang/String;
122+
}
123+
124+
public final class sk/ainet/compile/c/templates/HeaderTemplate {
125+
public static final field INSTANCE Lsk/ainet/compile/c/templates/HeaderTemplate;
126+
public final fun generate (Ljava/lang/String;[I[ILsk/ainet/compile/c/MemoryLayout;)Ljava/lang/String;
127+
}
128+
129+
public final class sk/ainet/compile/c/templates/SourceTemplate {
130+
public static final field INSTANCE Lsk/ainet/compile/c/templates/SourceTemplate;
131+
public final fun generate (Ljava/lang/String;Ljava/util/List;Ljava/util/List;)Ljava/lang/String;
132+
}
133+

0 commit comments

Comments
 (0)