394 lines
7.7 KiB
Go
394 lines
7.7 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|
// SPDX-License-Identifier: MIT
|
|
|
|
package linalg
|
|
|
|
import (
|
|
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
|
)
|
|
|
|
// Pipeline is a thin eager-evaluation wrapper around the package-level
|
|
// functions. Each step allocates a new array; the pipeline is
|
|
// just sugar over the same functions. It exists so callers can read
|
|
// sequential transformations as a single expression without losing the
|
|
// package-level style:
|
|
//
|
|
// out, err := Pipe(a).
|
|
// AddF(1).
|
|
// Sqrt().
|
|
// Result()
|
|
//
|
|
// The pipeline does not provide automatic error propagation: a step
|
|
// that returns an error is recorded, and subsequent steps become
|
|
// no-ops. Call Result() to retrieve both the final array and the first
|
|
// error observed (or nil).
|
|
//
|
|
// There is no lazy graph, no autograd, no backprop: just a chain of
|
|
// eager calls with the error short-circuit on top.
|
|
|
|
// Pipeline holds the current array under transformation and the first
|
|
// error encountered during the chain.
|
|
type Pipeline struct {
|
|
current *core.Array
|
|
err error
|
|
}
|
|
|
|
// Pipe starts a pipeline from a.
|
|
func Pipe(a *core.Array) *Pipeline {
|
|
return &Pipeline{current: a}
|
|
}
|
|
|
|
// Result returns the current array and the first error observed during
|
|
// the chain. If an error occurred earlier, current is the array that
|
|
// triggered it.
|
|
func (p *Pipeline) Result() (*core.Array, error) {
|
|
return p.current, p.err
|
|
}
|
|
|
|
// shortCircuit returns true when the pipeline already has an error:
|
|
// the caller becomes a no-op so subsequent steps stay composable.
|
|
func (p *Pipeline) shortCircuit() bool {
|
|
return p.err != nil
|
|
}
|
|
|
|
// assign updates the pipeline's current array under the wrapping rule
|
|
// (errors short-circuit, otherwise the new array is stored).
|
|
func (p *Pipeline) assign(out *core.Array, err error) {
|
|
if p.err != nil {
|
|
return
|
|
}
|
|
if err != nil {
|
|
p.err = err
|
|
return
|
|
}
|
|
p.current = out
|
|
}
|
|
|
|
// Binary array-array operations.
|
|
func (p *Pipeline) Add(b *core.Array) *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
out, err := core.Add(p.current, b)
|
|
p.assign(out, err)
|
|
return p
|
|
}
|
|
|
|
func (p *Pipeline) Sub(b *core.Array) *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
out, err := core.Sub(p.current, b)
|
|
p.assign(out, err)
|
|
return p
|
|
}
|
|
|
|
func (p *Pipeline) Mul(b *core.Array) *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
out, err := core.Mul(p.current, b)
|
|
p.assign(out, err)
|
|
return p
|
|
}
|
|
|
|
func (p *Pipeline) Div(b *core.Array) *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
out, err := core.Div(p.current, b)
|
|
p.assign(out, err)
|
|
return p
|
|
}
|
|
|
|
// Scalar arithmetic (the integer flavours only: float scalar variants
|
|
// would shadow Go's float64 promotion, so callers reach for AddF etc.
|
|
// explicitly).
|
|
func (p *Pipeline) AddI(v int64) *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
p.assign(core.AddI(p.current, v), nil)
|
|
return p
|
|
}
|
|
|
|
func (p *Pipeline) SubI(v int64) *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
p.assign(core.SubI(p.current, v), nil)
|
|
return p
|
|
}
|
|
|
|
func (p *Pipeline) MulI(v int64) *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
p.assign(core.MulI(p.current, v), nil)
|
|
return p
|
|
}
|
|
|
|
func (p *Pipeline) AddF(v float64) *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
p.assign(core.AddF(p.current, v), nil)
|
|
return p
|
|
}
|
|
|
|
func (p *Pipeline) SubF(v float64) *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
p.assign(core.SubF(p.current, v), nil)
|
|
return p
|
|
}
|
|
|
|
func (p *Pipeline) MulF(v float64) *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
p.assign(core.MulF(p.current, v), nil)
|
|
return p
|
|
}
|
|
|
|
// Reshape and transpose.
|
|
func (p *Pipeline) Reshape(shape ...int) *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
out, err := core.Reshape(p.current, shape...)
|
|
p.assign(out, err)
|
|
return p
|
|
}
|
|
|
|
func (p *Pipeline) Transpose() *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
p.assign(core.Transpose(p.current), nil)
|
|
return p
|
|
}
|
|
|
|
func (p *Pipeline) TransposeAxes(dims ...int) *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
out, err := core.TransposeAxes(p.current, dims...)
|
|
p.assign(out, err)
|
|
return p
|
|
}
|
|
|
|
func (p *Pipeline) Flatten(startDim, endDim int) *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
out, err := core.Flatten(p.current, startDim, endDim)
|
|
p.assign(out, err)
|
|
return p
|
|
}
|
|
|
|
func (p *Pipeline) Squeeze(dim int) *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
out, err := core.Squeeze(p.current, dim)
|
|
p.assign(out, err)
|
|
return p
|
|
}
|
|
|
|
func (p *Pipeline) Unsqueeze(dim int) *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
out, err := core.Unsqueeze(p.current, dim)
|
|
p.assign(out, err)
|
|
return p
|
|
}
|
|
|
|
// Reductions and extrema.
|
|
func (p *Pipeline) Maximum(b *core.Array) *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
out, err := core.Maximum(p.current, b)
|
|
p.assign(out, err)
|
|
return p
|
|
}
|
|
|
|
func (p *Pipeline) Minimum(b *core.Array) *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
out, err := core.Minimum(p.current, b)
|
|
p.assign(out, err)
|
|
return p
|
|
}
|
|
|
|
func (p *Pipeline) SumAxis(dim int) *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
out, err := core.SumAxis(p.current, dim)
|
|
p.assign(out, err)
|
|
return p
|
|
}
|
|
|
|
func (p *Pipeline) MeanAxis(dim int) *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
out, err := core.MeanAxis(p.current, dim)
|
|
p.assign(out, err)
|
|
return p
|
|
}
|
|
|
|
func (p *Pipeline) ClipI(lo, hi int64) *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
out, err := core.ClipI(p.current, lo, hi)
|
|
p.assign(out, err)
|
|
return p
|
|
}
|
|
|
|
func (p *Pipeline) ClipF(lo, hi float64) *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
out, err := core.ClipF(p.current, lo, hi)
|
|
p.assign(out, err)
|
|
return p
|
|
}
|
|
|
|
// Activations.
|
|
func (p *Pipeline) Tanh() *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
out, err := core.Tanh(p.current)
|
|
p.assign(out, err)
|
|
return p
|
|
}
|
|
|
|
func (p *Pipeline) Sigmoid() *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
out, err := core.Sigmoid(p.current)
|
|
p.assign(out, err)
|
|
return p
|
|
}
|
|
|
|
// Matrix and linear algebra.
|
|
func (p *Pipeline) MatMul2D(b *core.Array) *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
out, err := core.MatMul2D(p.current, b)
|
|
p.assign(out, err)
|
|
return p
|
|
}
|
|
|
|
func (p *Pipeline) Inv() *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
out, err := Inv(p.current)
|
|
p.assign(out, err)
|
|
return p
|
|
}
|
|
|
|
func (p *Pipeline) Solve(b *core.Array) *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
out, err := Solve(p.current, b)
|
|
p.assign(out, err)
|
|
return p
|
|
}
|
|
|
|
// Math functions.
|
|
func (p *Pipeline) Abs() *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
p.assign(core.Abs(p.current), nil)
|
|
return p
|
|
}
|
|
|
|
func (p *Pipeline) Neg() *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
p.assign(core.MulI(p.current, -1), nil)
|
|
return p
|
|
}
|
|
|
|
func (p *Pipeline) Sqrt() *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
out, err := core.Sqrt(p.current)
|
|
p.assign(out, err)
|
|
return p
|
|
}
|
|
|
|
func (p *Pipeline) Exp() *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
out, err := core.Exp(p.current)
|
|
p.assign(out, err)
|
|
return p
|
|
}
|
|
|
|
func (p *Pipeline) Log() *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
out, err := core.Log(p.current)
|
|
p.assign(out, err)
|
|
return p
|
|
}
|
|
|
|
func (p *Pipeline) Floor() *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
out, err := core.Floor(p.current)
|
|
p.assign(out, err)
|
|
return p
|
|
}
|
|
|
|
// Reductions and extrema.
|
|
func (p *Pipeline) ArgMaxAxis(dim int) *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
out, err := core.ArgMaxAxis(p.current, dim)
|
|
p.assign(out, err)
|
|
return p
|
|
}
|
|
|
|
func (p *Pipeline) ArgMinAxis(dim int) *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
out, err := core.ArgMinAxis(p.current, dim)
|
|
p.assign(out, err)
|
|
return p
|
|
}
|
|
|
|
// TopK keeps the top-k values along dim; the matching indices are
|
|
// discarded; call the package-level TopK directly when the indices
|
|
// matter.
|
|
func (p *Pipeline) TopK(k int, dim int) *Pipeline {
|
|
if p.shortCircuit() {
|
|
return p
|
|
}
|
|
values, _, err := core.TopK(p.current, k, dim)
|
|
p.assign(values, err)
|
|
return p
|
|
}
|