35 lines
1.2 KiB
Go
35 lines
1.2 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|||
|
|
// SPDX-License-Identifier: MIT
|
||
|
|
|
||
|
|
package core
|
||
|
|
|
||
|
|
import "testing"
|
||
|
|
|
||
|
|
// TestWhereHalfCanonicalisesNaN pins the float16 fast path of Where:
|
||
|
|
// every selected half value narrows through HalfFromFloat64 on the way
|
||
|
|
// out, so a NaN loses its payload and lands as the canonical 0x7E00 with
|
||
|
|
// its sign kept, while every finite value survives the round trip bit
|
||
|
|
// for bit. A payload copy would leak the incoming patterns 0x7C01 and
|
||
|
|
// 0x7D00 instead.
|
||
|
|
func TestWhereHalfCanonicalisesNaN(t *testing.T) {
|
||
|
|
cond := mustFromInts(t, []int64{1, 0, 1, 0}, 4)
|
||
|
|
// x: a signalling NaN with payload, 2, a negative signalling NaN, and
|
||
|
|
// a slot the condition does not pick. y: 4, 3, 4 and -1.
|
||
|
|
x := mustFromHalves(t, []uint16{0x7C01, 0x4000, 0xFE01, 0x3C00}, 4)
|
||
|
|
y := mustFromHalves(t, []uint16{0x4400, 0x4200, 0x4400, 0xBC00}, 4)
|
||
|
|
out, err := Where(cond, x, y)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("Where: %v", err)
|
||
|
|
}
|
||
|
|
if out.Dtype() != Float16 {
|
||
|
|
t.Fatalf("Where answered dtype %s, want float16", out.Dtype())
|
||
|
|
}
|
||
|
|
want := []uint16{0x7E00, 0x4200, 0xFE00, 0xBC00}
|
||
|
|
got := out.RawHalves()
|
||
|
|
for i, w := range want {
|
||
|
|
if got[i] != w {
|
||
|
|
t.Fatalf("element %d = %#04x, want %#04x", i, got[i], w)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|