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
4 changes: 4 additions & 0 deletions vecxt/src-js/doublematrix.scala
Original file line number Diff line number Diff line change
Expand Up @@ -303,6 +303,10 @@ object JsDoubleMatrix:
out
end *

/** Per-column fused multiply-add: `out(i, j) = m(i, j) * multiply(j) + add(j)`, as a fresh dense column-major
* matrix. See `JvmDoubleMatrix.fmaCols`; this platform uses the shared scalar loop.
*/
def fmaCols(multiply: Array[Double], add: Array[Double]): Matrix[Double] = FmaCols.loop(m, multiply, add)
end extension

end JsDoubleMatrix
Expand Down
48 changes: 48 additions & 0 deletions vecxt/src-jvm/doublematrix.scala
Original file line number Diff line number Diff line change
Expand Up @@ -384,6 +384,54 @@ object JvmDoubleMatrix:

end +=

/** Per-column fused multiply-add: `out(i, j) = m(i, j) * multiply(j) + add(j)`, as a fresh dense column-major
* matrix.
*
* A column scale and a column broadcast-add in one pass: `m` is read once and the result written once, where
* `copy`, `+= arr` and a column scale would each make a full pass. Reference BLAS has no equivalent - `dscal` per
* column scales but cannot add, and `A * diag(x)` routines (`dgmm`) likewise stop at the multiply.
*
* When `rowStride == 1` each column is a contiguous run and is processed with `DoubleVector.fma` against two
* broadcast lanes; any other layout (row-major, transposed or doubly-strided views) falls back to the shared
* scalar loop in [[FmaCols.loop]]. The input is never modified.
*
* @param multiply
* one multiplier per column; must have length `m.cols`
* @param add
* one addend per column; must have length `m.cols`
*/
def fmaCols(multiply: Array[Double], add: Array[Double]): Matrix[Double] =
if m.rowStride != 1 then FmaCols.loop(m, multiply, add)
else
FmaCols.check(m, multiply, add)
val spd = doublearrays.spd
val spdl = doublearrays.spdl
val rows = m.rows
val bound = spd.loopBound(rows)
val out = new Array[Double](m.numel)
var j = 0
while j < m.cols do
val src = m.offset + j * m.colStride
val dst = j * rows
val mulS = multiply(j)
val addS = add(j)
val mul = DoubleVector.broadcast(spd, mulS)
val ad = DoubleVector.broadcast(spd, addS)
var i = 0
while i < bound do
DoubleVector.fromArray(spd, m.raw, src + i).fma(mul, ad).intoArray(out, dst + i)
i += spdl
end while
while i < rows do
out(dst + i) = Math.fma(m.raw(src + i), mulS, addS)
i += 1
end while
j += 1
end while
Matrix[Double](out, m.rows, m.cols)
end if
end fmaCols

def +=(n: Double): Unit =

if m.hasSimpleContiguousMemoryLayout then vecxt.doublearrays.+=(m.raw)(n)
Expand Down
5 changes: 5 additions & 0 deletions vecxt/src-native/doublematrix_native.scala
Original file line number Diff line number Diff line change
Expand Up @@ -308,6 +308,11 @@ object NativeDoubleMatrix:
m.*=(vec, out, alpha, 0.0)
out
end *

/** Per-column fused multiply-add: `out(i, j) = m(i, j) * multiply(j) + add(j)`, as a fresh dense column-major
* matrix. See `JvmDoubleMatrix.fmaCols`; this platform uses the shared scalar loop.
*/
def fmaCols(multiply: Array[Double], add: Array[Double]): Matrix[Double] = FmaCols.loop(m, multiply, add)
end extension

end NativeDoubleMatrix
Expand Down
47 changes: 47 additions & 0 deletions vecxt/src/fmaCols.scala
Original file line number Diff line number Diff line change
@@ -0,0 +1,47 @@
package vecxt

import vecxt.matrix.*

/** Shared pieces of `fmaCols`: `out(i, j) = m(i, j) * multiply(j) + add(j)`.
*
* Each platform exposes the public `fmaCols` extension (`JvmDoubleMatrix`, `JsDoubleMatrix`, `NativeDoubleMatrix`). JS
* and Native delegate straight to [[loop]]; the JVM uses a SIMD path when rows are contiguous and falls back to
* [[loop]] otherwise. Deliberately not exported from `vecxt.all` - it is the kernel, not the API.
*/
private[vecxt] object FmaCols:

/** Throws unless both vectors have one entry per column of `m`. Message built only on the failing path. */
inline def check(m: Matrix[Double], multiply: Array[Double], add: Array[Double]): Unit =
if multiply.length != m.cols then
throw new IllegalArgumentException(s"multiply length ${multiply.length} != expected ${m.cols}")
end if
if add.length != m.cols then throw new IllegalArgumentException(s"add length ${add.length} != expected ${m.cols}")
end if
end check

/** Layout-agnostic kernel: reads through the strides, writes a fresh dense column-major matrix.
*
* Columns outer, rows inner, so the write is always sequential. Uses a plain `x * mul + ad` rather than `Math.fma`,
* which is emulated (slowly) on Scala.js; results may therefore differ from the JVM SIMD path in the last bit.
*/
def loop(m: Matrix[Double], multiply: Array[Double], add: Array[Double]): Matrix[Double] =
check(m, multiply, add)
val rows = m.rows
val out = new Array[Double](m.numel)
var j = 0
while j < m.cols do
val src = m.offset + j * m.colStride
val dst = j * rows
val mul = multiply(j)
val ad = add(j)
var i = 0
while i < rows do
out(dst + i) = m.raw(src + i * m.rowStride) * mul + ad
i += 1
end while
j += 1
end while
Matrix[Double](out, m.rows, m.cols)
end loop

end FmaCols
80 changes: 80 additions & 0 deletions vecxt/test/src/fmaCols.test.scala
Original file line number Diff line number Diff line change
@@ -0,0 +1,80 @@
package vecxt

import all.*

/** `fmaCols` is `out(i, j) = m(i, j) * multiply(j) + add(j)`. The JVM has a SIMD path for `rowStride == 1` and falls
* back to the shared scalar loop otherwise; JS and Native always use the loop. Every case is checked against a naive
* reference read through `m(i, j)`, so the layouts that pick different branches are held to the same answer.
*/
class FmaColsSuite extends munit.FunSuite:

private def reference(m: Matrix[Double], multiply: Array[Double], add: Array[Double]): Matrix[Double] =
val out = Array.ofDim[Double](m.rows * m.cols)
for j <- 0 until m.cols; i <- 0 until m.rows do out(i + j * m.rows) = m(i, j) * multiply(j) + add(j)
end for
Matrix[Double](out, m.rows, m.cols)
end reference

private def assertDenseColMajor(m: Matrix[Double])(using munit.Location): Unit =
assertEquals((m.rowStride, m.colStride, m.offset), (1, m.rows, 0))

private def check(m: Matrix[Double], multiply: Array[Double], add: Array[Double])(using munit.Location): Unit =
val out = m.fmaCols(multiply, add)
assertDenseColMajor(out)
assertMatrixEquals(out, reference(m, multiply, add))
end check

test("small hand computed example"):
val m = Matrix.fromRows[Double](Array(1.0, 2.0), Array(3.0, 4.0))
val out = m.fmaCols(Array(10.0, -1.0), Array(0.5, 100.0))
assertMatrixEquals(out, Matrix.fromRows[Double](Array(10.5, 98.0), Array(30.5, 96.0)))

test("dense column-major, row count not a multiple of common SIMD widths"):
val m = Matrix[Double](Array.tabulate(37 * 3)(i => i * 0.25 - 3.0), 37, 3)
check(m, Array(2.0, -0.5, 0.0), Array(1.0, 7.0, -3.0))

test("dense column-major, large enough for many SIMD iterations per column"):
val m = Matrix[Double](Array.tabulate(1000 * 4)(i => math.sin(i.toDouble)), 1000, 4)
check(m, Array(1.5, -2.0, 3.25, 1e6), Array(-1.0, 0.0, 2.0, 1e-3))

test("row-major input"):
val m = Matrix[Double](Array.tabulate(5 * 3)(_.toDouble), 5, 3, 3, 1, 0)
assertEquals(m.rowStride, 3)
check(m, Array(1.0, 2.0, 3.0), Array(-1.0, -2.0, -3.0))

test("padded column stride (rowStride 1, not contiguous)"):
// rows=2, cols=2, colStride=3: columns are raw(0..1) and raw(3..4), raw(2)/raw(5) are padding
val m = Matrix[Double](Array(1.0, 2.0, 99.0, 3.0, 4.0, 99.0), 2, 2, 1, 3, 0)
assert(!m.hasSimpleContiguousMemoryLayout)
val out = m.fmaCols(Array(2.0, 3.0), Array(1.0, 1.0))
assertMatrixEquals(out, Matrix.fromRows[Double](Array(3.0, 10.0), Array(5.0, 13.0)))

test("submatrix view with offset"):
val base = Matrix[Double](Array.tabulate(20 * 6)(_.toDouble), 20, 6)
// rows 3..17 of columns 2..4 of a 20 x 6 column-major matrix
val view = Matrix[Double](base.raw, 15, 3, 1, 20, 3 + 2 * 20)
assertEquals(view(0, 0), 43.0)
check(view, Array(-1.0, 0.5, 2.0), Array(10.0, 20.0, 30.0))

test("input is not modified"):
val raw = Array.tabulate(9 * 2)(_.toDouble)
val m = Matrix[Double](raw.clone, 9, 2)
m.fmaCols(Array(3.0, 4.0), Array(1.0, 2.0))
assertVecEquals(m.raw, raw)

test("empty matrices"):
val noRows = Matrix[Double](Array.empty[Double], 0, 3)
val out = noRows.fmaCols(Array(1.0, 2.0, 3.0), Array(1.0, 2.0, 3.0))
assertEquals((out.rows, out.cols), (0, 3))
val noCols = Matrix[Double](Array.empty[Double], 4, 0)
val out2 = noCols.fmaCols(Array.empty[Double], Array.empty[Double])
assertEquals((out2.rows, out2.cols), (4, 0))

test("wrong vector lengths throw"):
val m = Matrix[Double](Array.fill(6)(1.0), 3, 2)
intercept[IllegalArgumentException](m.fmaCols(Array(1.0), Array(1.0, 2.0)))
intercept[IllegalArgumentException](m.fmaCols(Array(1.0, 2.0), Array(1.0, 2.0, 3.0)))
val rowMajor = Matrix[Double](Array.fill(6)(1.0), 3, 2, 2, 1, 0)
intercept[IllegalArgumentException](rowMajor.fmaCols(Array(1.0), Array(1.0, 2.0)))

end FmaColsSuite
75 changes: 75 additions & 0 deletions vecxt_re/src/CurrencyValue.scala
Original file line number Diff line number Diff line change
@@ -0,0 +1,75 @@
package vecxt_re

enum Ccy derives CanEqual:
case USD, EUR, GBP, CHF, JPY, AUD, CAD, NZD
end Ccy

/** The scale a monetary amount is quoted in. `multiplier` is the number of major currency units in one of this unit. */
enum AmountUnit(val multiplier: Long):
case One extends AmountUnit(1L)
case K extends AmountUnit(1_000L)
case Mn extends AmountUnit(1_000_000L)
case Bn extends AmountUnit(1_000_000_000L)
end AmountUnit

/** This should _NOT_ be used for actual cash values. We use double here because this intended for financial simulation
* where errors of <0.01 are meaningless. An absolute monetary quantity denominated in a single currency, e.g. a
* notional, limit, premium or loss.
*
* Amounts in different currencies are never implicitly combined: arithmetic between currencies fails, and conversion
* requires an explicit FX rate. Scaling by a [[Rel]] yields a CurrencyAmount (premium = notional × coupon); dividing
* two amounts in the same currency yields a [[Rel]].
*
* Stored as a Double in major currency units; [[AmountUnit]] only affects construction and display.
*/
final case class CurrencyAmount private (amount: Double, ccy: Ccy) derives CanEqual:

/** The amount expressed in unit `u`, e.g. `CurrencyAmount(2.5, Mn, USD).in(K) == 2500.0` */
def in(u: AmountUnit): Double = amount / u.multiplier

def +(y: CurrencyAmount): CurrencyAmount =
sameCcy(y); new CurrencyAmount(amount + y.amount, ccy)
end +

def -(y: CurrencyAmount): CurrencyAmount =
sameCcy(y); new CurrencyAmount(amount - y.amount, ccy)
end -

def unary_- : CurrencyAmount = new CurrencyAmount(-amount, ccy)

/** premium = notional * coupon */
def *(r: Rel): CurrencyAmount =
new CurrencyAmount(amount * r.canonical, ccy)

def *(k: Double): CurrencyAmount =
new CurrencyAmount(amount * k, ccy)

/** Ratio of two same-currency amounts, e.g. a loss ratio. */
def /(y: CurrencyAmount): Rel =
sameCcy(y); Rel(amount / y.amount, PriceUnit.One)
end /

def approxEq(y: CurrencyAmount, tol: Double = 1e-6): Boolean =
sameCcy(y); math.abs(amount - y.amount) <= tol
end approxEq

override def toString: String = f"$amount%.2f $ccy"

private def sameCcy(y: CurrencyAmount): Unit =
require(ccy == y.ccy, s"Currency mismatch: $ccy vs ${y.ccy}")
end CurrencyAmount

object CurrencyAmount:
/** e.g. `CurrencyAmount(2.5, AmountUnit.Mn, Ccy.USD)` is 2,500,000 USD */
def apply(amount: Double, unit: AmountUnit, ccy: Ccy): CurrencyAmount =
new CurrencyAmount(amount * unit.multiplier, ccy)

def zero(ccy: Ccy): CurrencyAmount = new CurrencyAmount(0.0, ccy)

given Ordering[CurrencyAmount] with
def compare(a: CurrencyAmount, b: CurrencyAmount): Int =
require(a.ccy == b.ccy, s"Cannot order ${a.ccy} vs ${b.ccy}")
java.lang.Double.compare(a.amount, b.amount)
end compare
end given
end CurrencyAmount
55 changes: 55 additions & 0 deletions vecxt_re/src/NamedRecord.scala
Original file line number Diff line number Diff line change
@@ -0,0 +1,55 @@
package vecxt_re

import scala.NamedTuple.{AnyNamedTuple, DropNames, NamedTuple, Names}
import scala.compiletime.{constValue, constValueTuple, erasedValue, error, summonFrom}

/** Compile-time helpers for accepting a named tuple by field name rather than by field position.
*
* A caller may supply `(id = "1", price = p, notional = n, name = "A")` where the required shape is
* `(price: Price, notional: CurrencyAmount, name: String, id: String)`: field order is irrelevant and extra fields are
* ignored, but every required field must be present with a conforming type or compilation fails.
*/
object NamedRecord:

/** Marker produced by [[Lookup]] when a field name is absent. */
sealed trait Missing

/** The value type of field `K` in the named tuple with names `N` and values `V`, or [[Missing]]. */
type Lookup[N <: Tuple, V <: Tuple, K] = (N, V) match
case (K *: _, v *: _) => v
case (_ *: ns, _ *: vs) => Lookup[ns, vs, K]
case (EmptyTuple, EmptyTuple) => Missing

/** Fails compilation unless the named tuple `(N, V)` has every field of `Required` with a conforming type.
*
* An unconstrained `N` (inferred as `Tuple` when passing `Nil`) is accepted, as there are no values to check.
*/
inline def check[N <: Tuple, V <: Tuple, Required <: AnyNamedTuple]: Unit =
summonFrom {
case _: (N =:= Tuple) => ()
case _ => checkFields[N, V, Names[Required], DropNames[Required]]
}

private inline def checkFields[N <: Tuple, V <: Tuple, RN <: Tuple, RV <: Tuple]: Unit =
inline erasedValue[RN] match
case _: EmptyTuple => ()
case _: (k *: rns) =>
inline erasedValue[RV] match
case _: (t *: rvs) =>
inline erasedValue[Lookup[N, V, k]] match
case _: Missing => error("Missing field: " + constValue[k & String])
case _: t => checkFields[N, V, rns, rvs]
case _ => error("Field has the wrong type: " + constValue[k & String])

/** The field names of `N` as runtime strings; empty for an unconstrained `N` (see [[check]]). */
inline def names[N <: Tuple]: List[String] =
summonFrom {
case _: (N =:= Tuple) => Nil
case _ => constValueTuple[N].toList.asInstanceOf[List[String]]
}

/** Zips pre-computed field `names` with the values of `nt` into a name -> value map. */
def toMap[N <: Tuple, V <: Tuple](names: List[String], nt: NamedTuple[N, V]): Map[String, Any] =
names.iterator.zip(nt.toTuple.productIterator).toMap

end NamedRecord
Loading
Loading