Skip to content

Commit b121652

Browse files
authored
Merge pull request #207 from T45K/add_norm_for_D1Array
Add `norm` API for vectors
2 parents 79dedd1 + 6b3071f commit b121652

3 files changed

Lines changed: 38 additions & 0 deletions

File tree

  • multik-core/src/commonMain/kotlin/org/jetbrains/kotlinx/multik/api/linalg
  • multik-kotlin/src/commonTest/kotlin/org/jetbrains/kotlinx/multik_kotlin/linAlg
  • multik-openblas/src/commonTest/kotlin/org/jetbrains/kotlinx/multik/openblas/linalg

multik-core/src/commonMain/kotlin/org/jetbrains/kotlinx/multik/api/linalg/norm.kt

Lines changed: 18 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -4,16 +4,34 @@
44

55
package org.jetbrains.kotlinx.multik.api.linalg
66

7+
import org.jetbrains.kotlinx.multik.api.mk
8+
import org.jetbrains.kotlinx.multik.api.zeros
9+
import org.jetbrains.kotlinx.multik.ndarray.data.D1
710
import org.jetbrains.kotlinx.multik.ndarray.data.D2
811
import org.jetbrains.kotlinx.multik.ndarray.data.MultiArray
12+
import org.jetbrains.kotlinx.multik.ndarray.operations.stack
913
import kotlin.jvm.JvmName
1014

15+
/**
16+
* Returns norm of float vector
17+
*/
18+
@JvmName("normFV")
19+
public fun LinAlg.norm(mat: MultiArray<Float, D1>, norm: Norm = Norm.Fro): Float =
20+
this.linAlgEx.normF(mk.stack(mat, mk.zeros(mat.size)), norm)
21+
1122
/**
1223
* Returns norm of float matrix
1324
*/
1425
@JvmName("normF")
1526
public fun LinAlg.norm(mat: MultiArray<Float, D2>, norm: Norm = Norm.Fro): Float = this.linAlgEx.normF(mat, norm)
1627

28+
/**
29+
* Returns norm of double vector
30+
*/
31+
@JvmName("normDV")
32+
public fun LinAlg.norm(mat: MultiArray<Double, D1>, norm: Norm = Norm.Fro): Double =
33+
this.linAlgEx.norm(mk.stack(mat, mk.zeros(mat.size)), norm)
34+
1735
/**
1836
* Returns norm of double matrix
1937
*/

multik-kotlin/src/commonTest/kotlin/org/jetbrains/kotlinx/multik_kotlin/linAlg/KELinAlgTest.kt

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -492,6 +492,16 @@ class KELinAlgTest {
492492
}
493493

494494
}
495+
496+
@Test
497+
fun compute_norm_for_vector() {
498+
val vector = mk.ndarray(mk[1.1, 0.0, 3.2, 2.3, 5.0])
499+
500+
assertEquals(6.460650122085238, mk.linalg.norm(vector, Norm.Fro))
501+
assertEquals(11.600000000000001, mk.linalg.norm(vector, Norm.Inf))
502+
assertEquals(5.0, mk.linalg.norm(vector, Norm.N1))
503+
assertEquals(5.0, mk.linalg.norm(vector, Norm.Max))
504+
}
495505
}
496506

497507

multik-openblas/src/commonTest/kotlin/org/jetbrains/kotlinx/multik/openblas/linalg/NativeLinAlgTest.kt

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -324,4 +324,14 @@ class NativeLinAlgTest {
324324
assertFloatingNumber(7.0, NativeLinAlg.norm(b, Norm.N1))
325325
assertFloatingNumber(4.0, NativeLinAlg.norm(b, Norm.Max))
326326
}
327+
328+
@Test
329+
fun `compute norm for vector`() {
330+
val vector = mk.ndarray(mk[1.1, 0.0, 3.2, 2.3, 5.0])
331+
332+
assertEquals(6.460650122085238, mk.linalg.norm(vector, Norm.Fro))
333+
assertEquals(11.600000000000001, mk.linalg.norm(vector, Norm.Inf))
334+
assertEquals(5.0, mk.linalg.norm(vector, Norm.N1))
335+
assertEquals(5.0, mk.linalg.norm(vector, Norm.Max))
336+
}
327337
}

0 commit comments

Comments
 (0)