Files

214 lines
7.0 KiB
Go
Raw Permalink Normal View History

2026-09-03 10:00:00 +02:00
// 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
}