diff --git a/skainet-lang/skainet-lang-core/api/android/skainet-lang-core.api b/skainet-lang/skainet-lang-core/api/android/skainet-lang-core.api index 79f3a5ad..369d614a 100644 --- a/skainet-lang/skainet-lang-core/api/android/skainet-lang-core.api +++ b/skainet-lang/skainet-lang-core/api/android/skainet-lang-core.api @@ -754,7 +754,7 @@ public abstract class sk/ainet/lang/nn/InternalMixedPrecisionModule : sk/ainet/l protected abstract fun forwardImpl (Lsk/ainet/lang/tensor/Tensor;)Lsk/ainet/lang/tensor/Tensor; } -public final class sk/ainet/lang/nn/Linear : sk/ainet/lang/nn/Module, sk/ainet/lang/nn/topology/ModuleParameters { +public class sk/ainet/lang/nn/Linear : sk/ainet/lang/nn/Module, sk/ainet/lang/nn/topology/ModuleParameters { public fun (IILjava/lang/String;Lsk/ainet/lang/tensor/Tensor;)V public fun (IILjava/lang/String;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;)V public fun (IILjava/lang/String;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Z)V @@ -764,6 +764,7 @@ public final class sk/ainet/lang/nn/Linear : sk/ainet/lang/nn/Module, sk/ainet/l public fun getName ()Ljava/lang/String; public fun getParams ()Ljava/util/List; public final fun getTrainable ()Z + protected fun onForward (Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/context/ExecutionContext;)Lsk/ainet/lang/tensor/Tensor; } public final class sk/ainet/lang/nn/MaxPool2d : sk/ainet/lang/nn/Module { diff --git a/skainet-lang/skainet-lang-core/api/jvm/skainet-lang-core.api b/skainet-lang/skainet-lang-core/api/jvm/skainet-lang-core.api index 2b189600..338bceb0 100644 --- a/skainet-lang/skainet-lang-core/api/jvm/skainet-lang-core.api +++ b/skainet-lang/skainet-lang-core/api/jvm/skainet-lang-core.api @@ -989,7 +989,7 @@ public final class sk/ainet/lang/nn/LayerScale : sk/ainet/lang/nn/Module, sk/ain public fun getParams ()Ljava/util/List; } -public final class sk/ainet/lang/nn/Linear : sk/ainet/lang/nn/Module, sk/ainet/lang/nn/topology/ModuleParameters { +public class sk/ainet/lang/nn/Linear : sk/ainet/lang/nn/Module, sk/ainet/lang/nn/topology/ModuleParameters { public fun (IILjava/lang/String;Lsk/ainet/lang/tensor/Tensor;)V public fun (IILjava/lang/String;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;)V public fun (IILjava/lang/String;Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/lang/tensor/Tensor;Z)V @@ -999,6 +999,7 @@ public final class sk/ainet/lang/nn/Linear : sk/ainet/lang/nn/Module, sk/ainet/l public fun getName ()Ljava/lang/String; public fun getParams ()Ljava/util/List; public final fun getTrainable ()Z + protected fun onForward (Lsk/ainet/lang/tensor/Tensor;Lsk/ainet/context/ExecutionContext;)Lsk/ainet/lang/tensor/Tensor; } public final class sk/ainet/lang/nn/Lstm : sk/ainet/lang/nn/Module, sk/ainet/lang/nn/topology/ModuleParameters { diff --git a/skainet-lang/skainet-lang-core/src/commonMain/kotlin/sk/ainet/lang/nn/Linear.kt b/skainet-lang/skainet-lang-core/src/commonMain/kotlin/sk/ainet/lang/nn/Linear.kt index e55cbfa1..34154fb1 100644 --- a/skainet-lang/skainet-lang-core/src/commonMain/kotlin/sk/ainet/lang/nn/Linear.kt +++ b/skainet-lang/skainet-lang-core/src/commonMain/kotlin/sk/ainet/lang/nn/Linear.kt @@ -19,6 +19,9 @@ import sk.ainet.lang.nn.topology.weights * heads. A bias-less layer registers only the weight parameter, so parameter * counts and checkpoints match architectures defined without bias. * + * The class is `open` so adapter-style layers (e.g. LoRA) can subclass it and + * augment [onForward] or [params] while reusing the base projection. + * * @param inFeatures Number of input features * @param outFeatures Number of output features * @param name Name of the module @@ -26,7 +29,7 @@ import sk.ainet.lang.nn.topology.weights * @param initBias Initial bias, or `null` for a layer without bias */ -public class Linear @kotlin.jvm.JvmOverloads constructor( +public open class Linear @kotlin.jvm.JvmOverloads constructor( inFeatures: Int, outFeatures: Int, override val name: String = "Linear",