214 lines
7.0 KiB
Go
214 lines
7.0 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|||
|
|
// SPDX-License-Identifier: MIT
|
||
|
|
|
||
|
|
package core
|
||
|
|
|
||
|
|
// Element-wise extrema and clipping, plus the Go 1.27 generic
|
||
|
|
// payload accessor. Minimum and Maximum follow IEEE NaN propagation: a
|
||
|
|
// NaN operand poisons the result element; complex arrays have no
|
||
|
|
// ordering and error.
|
||
|
|
|
||
|
|
// Minimum returns the element-wise smaller of two same-shape arrays; NaN
|
||
|
|
// propagates.
|
||
|
|
func Minimum(a, b *Array) (*Array, error) {
|
||
|
|
return elementwise(a, b, "Minimum",
|
||
|
|
func(x, y int64) int64 { return min(x, y) },
|
||
|
|
func(x, y float64) float64 { return min(x, y) },
|
||
|
|
nil)
|
||
|
|
}
|
||
|
|
|
||
|
|
// Maximum returns the element-wise larger of two same-shape arrays; NaN
|
||
|
|
// propagates.
|
||
|
|
func Maximum(a, b *Array) (*Array, error) {
|
||
|
|
return elementwise(a, b, "Maximum",
|
||
|
|
func(x, y int64) int64 { return max(x, y) },
|
||
|
|
func(x, y float64) float64 { return max(x, y) },
|
||
|
|
nil)
|
||
|
|
}
|
||
|
|
|
||
|
|
// ClipI clamps each element of a real array into [lo, hi]; the dtype
|
||
|
|
// keeps its kind. lo > hi is an error.
|
||
|
|
func ClipI(a *Array, lo, hi int64) (*Array, error) {
|
||
|
|
if lo > hi {
|
||
|
|
return nil, errf("ClipI: lo must be at most hi, got %d and %d", lo, hi)
|
||
|
|
}
|
||
|
|
switch a.dt {
|
||
|
|
case Int:
|
||
|
|
out := &Array{shape: a.Shape(), dt: Int, ints: make([]int64, a.Len())}
|
||
|
|
parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) {
|
||
|
|
as, os := a.ints[s:e], out.ints[s:e]
|
||
|
|
for i := range os {
|
||
|
|
os[i] = min(max(as[i], lo), hi)
|
||
|
|
}
|
||
|
|
})
|
||
|
|
return out, nil
|
||
|
|
case Float16:
|
||
|
|
flo, fhi := float64(lo), float64(hi)
|
||
|
|
out := &Array{shape: a.Shape(), dt: Float16, halves: make([]uint16, a.Len())}
|
||
|
|
if a.strides == nil {
|
||
|
|
parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) {
|
||
|
|
as, os := a.halves[s:e], out.halves[s:e]
|
||
|
|
for i := range os {
|
||
|
|
// Clamped in float64, where the half values compare
|
||
|
|
// exactly, then narrowed once: half bit patterns do not
|
||
|
|
// order as uint16.
|
||
|
|
os[i] = HalfFromFloat64(min(max(HalfToFloat64(as[i]), flo), fhi))
|
||
|
|
}
|
||
|
|
})
|
||
|
|
return out, nil
|
||
|
|
}
|
||
|
|
parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) {
|
||
|
|
for i := s; i < e; i++ {
|
||
|
|
out.halves[i] = HalfFromFloat64(min(max(a.halfAt(i), flo), fhi))
|
||
|
|
}
|
||
|
|
})
|
||
|
|
return out, nil
|
||
|
|
case Float32:
|
||
|
|
flo, fhi := float32(lo), float32(hi)
|
||
|
|
out := &Array{shape: a.Shape(), dt: Float32, floats32: make([]float32, a.Len())}
|
||
|
|
parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) {
|
||
|
|
as, os := a.floats32[s:e], out.floats32[s:e]
|
||
|
|
for i := range os {
|
||
|
|
os[i] = min(max(as[i], flo), fhi)
|
||
|
|
}
|
||
|
|
})
|
||
|
|
return out, nil
|
||
|
|
case Float:
|
||
|
|
flo, fhi := float64(lo), float64(hi)
|
||
|
|
out := &Array{shape: a.Shape(), dt: Float, floats: make([]float64, a.Len())}
|
||
|
|
parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) {
|
||
|
|
as, os := a.floats[s:e], out.floats[s:e]
|
||
|
|
for i := range os {
|
||
|
|
os[i] = min(max(as[i], flo), fhi)
|
||
|
|
}
|
||
|
|
})
|
||
|
|
return out, nil
|
||
|
|
case Complex:
|
||
|
|
return nil, errf("ClipI: complex arrays have no ordering")
|
||
|
|
default:
|
||
|
|
// Bool and the narrow integer widths carry no clamp kernel here;
|
||
|
|
// the message names the actual dtype rather than complex.
|
||
|
|
return nil, errf("ClipI: dtype %s is not supported; convert with Astype", a.dt)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// ClipF clamps each element of a real array into [lo, hi]; int arrays
|
||
|
|
// become float, float16 and float32 arrays keep their dtype with the
|
||
|
|
// clamp run in float64 and narrowed once. lo > hi is an error.
|
||
|
|
func ClipF(a *Array, lo, hi float64) (*Array, error) {
|
||
|
|
if lo > hi {
|
||
|
|
return nil, errf("ClipF: lo must be at most hi, got %v and %v", lo, hi)
|
||
|
|
}
|
||
|
|
if a.dt == Complex {
|
||
|
|
return nil, errf("ClipF: complex arrays have no ordering")
|
||
|
|
}
|
||
|
|
if narrowRefused(a.dt) {
|
||
|
|
// Bool and the narrow integer widths carry no clamp kernel here;
|
||
|
|
// the refusal names the dtype and the conversion.
|
||
|
|
return nil, errf("ClipF: dtype %s is not supported; convert with Astype", a.dt)
|
||
|
|
}
|
||
|
|
out := &Array{shape: a.Shape(), dt: a.dt}
|
||
|
|
if a.dt == Int {
|
||
|
|
out.dt = Float
|
||
|
|
}
|
||
|
|
out.alloc(a.Len())
|
||
|
|
switch a.dt {
|
||
|
|
case Float16:
|
||
|
|
if a.strides == nil {
|
||
|
|
parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) {
|
||
|
|
as, os := a.halves[s:e], out.halves[s:e]
|
||
|
|
for i := range os {
|
||
|
|
os[i] = HalfFromFloat64(min(max(HalfToFloat64(as[i]), lo), hi))
|
||
|
|
}
|
||
|
|
})
|
||
|
|
} else {
|
||
|
|
parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) {
|
||
|
|
for i := s; i < e; i++ {
|
||
|
|
out.halves[i] = HalfFromFloat64(min(max(a.halfAt(i), lo), hi))
|
||
|
|
}
|
||
|
|
})
|
||
|
|
}
|
||
|
|
case Float32:
|
||
|
|
flo, fhi := float32(lo), float32(hi)
|
||
|
|
parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) {
|
||
|
|
as, os := a.floats32[s:e], out.floats32[s:e]
|
||
|
|
for i := range os {
|
||
|
|
os[i] = min(max(as[i], flo), fhi)
|
||
|
|
}
|
||
|
|
})
|
||
|
|
case Int:
|
||
|
|
if a.strides == nil {
|
||
|
|
parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) {
|
||
|
|
as, os := a.ints[s:e], out.floats[s:e]
|
||
|
|
for i := range os {
|
||
|
|
os[i] = min(max(float64(as[i]), lo), hi)
|
||
|
|
}
|
||
|
|
})
|
||
|
|
} else {
|
||
|
|
parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) {
|
||
|
|
for i := s; i < e; i++ {
|
||
|
|
out.floats[i] = min(max(a.floatAt(i), lo), hi)
|
||
|
|
}
|
||
|
|
})
|
||
|
|
}
|
||
|
|
default:
|
||
|
|
parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) {
|
||
|
|
as, os := a.floats[s:e], out.floats[s:e]
|
||
|
|
for i := range os {
|
||
|
|
os[i] = min(max(as[i], lo), hi)
|
||
|
|
}
|
||
|
|
})
|
||
|
|
}
|
||
|
|
return out, nil
|
||
|
|
}
|
||
|
|
|
||
|
|
// Elements returns the flat payload converted to the caller's type, a
|
||
|
|
// Go 1.27 generic method. E may hold anything the payload ladder
|
||
|
|
// widens to: the ladder runs int64 ← float16 ← float32 ← float64 ←
|
||
|
|
// complex128, and requesting a rung wider than the array's converts
|
||
|
|
// while a narrower one is an error naming both dtypes. E = int64
|
||
|
|
// follows the rule IntAt carries: the whole integer class, bool
|
||
|
|
// included, widens exactly through intAt, and a float or complex
|
||
|
|
// source has no exact int64 image, so it keeps the narrowing refusal.
|
||
|
|
// The returned slice is a copy; changing it never reaches the array.
|
||
|
|
func (a *Array) Elements[E int64 | float32 | float64 | complex128]() ([]E, error) {
|
||
|
|
out := make([]E, a.Len())
|
||
|
|
var zero E
|
||
|
|
switch any(zero).(type) {
|
||
|
|
case complex128:
|
||
|
|
// Every dtype converts to complex.
|
||
|
|
for i := range a.Len() {
|
||
|
|
out[i] = any(a.complexAt(i)).(E)
|
||
|
|
}
|
||
|
|
case float64:
|
||
|
|
// int, float32 and float64 convert; complex is a narrowing.
|
||
|
|
if a.dt == Complex {
|
||
|
|
return nil, errf("Elements: cannot narrow %s to float64", a.dt)
|
||
|
|
}
|
||
|
|
for i := range a.Len() {
|
||
|
|
out[i] = any(a.floatAt(i)).(E)
|
||
|
|
}
|
||
|
|
case float32:
|
||
|
|
// int and float32 convert; float64 and complex are narrowings.
|
||
|
|
if a.dt == Float || a.dt == Complex {
|
||
|
|
return nil, errf("Elements: cannot narrow %s to float32", a.dt)
|
||
|
|
}
|
||
|
|
for i := range a.Len() {
|
||
|
|
out[i] = any(a.float32At(i)).(E)
|
||
|
|
}
|
||
|
|
default:
|
||
|
|
// int64: the integer class widens exactly through intAt, the
|
||
|
|
// same widening IntAt answers; float and complex sources are
|
||
|
|
// genuine narrowings and keep the refusal. intAt gathers a
|
||
|
|
// strided array through physIndex, exactly as the float32
|
||
|
|
// branch's accessor does, instead of reading the payload raw.
|
||
|
|
if !intClass(a.dt) {
|
||
|
|
return nil, errf("Elements: cannot narrow %s to int64", a.dt)
|
||
|
|
}
|
||
|
|
for i := range a.Len() {
|
||
|
|
out[i] = any(a.intAt(i)).(E)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
return out, nil
|
||
|
|
}
|