--wip-- [skip ci]

This commit is contained in:
2025-06-24 23:22:06 +01:00
parent 77c2610743
commit 3ca29db4c8
+287 -152
View File
@@ -4,12 +4,103 @@ use std::simd::{
StdFloat,
cmp::{SimdPartialEq, SimdPartialOrd},
num::SimdFloat,
u32x8,
};
use crate::interpreter::{
Mask, VALUE_0, VALUE_1, VALUE_2, VALUE_05, VALUE_M1, VALUE_NAN, VALUE_PI, Value,
};
/// A single choice made at a min/max node.
///
/// Explicitly stored in a `u8` so that this can be written by JIT functions,
/// which have no notion of Rust enums.
///
/// Note that this is a bitfield such that
/// ```rust
/// # use fidget::vm::Choice;
/// # assert!(
/// Choice::Both as u8 == Choice::Left as u8 | Choice::Right as u8
/// # );
/// ```
#[derive(Copy, Clone, Debug, Eq, PartialEq)]
#[repr(u8)]
pub enum Choice {
/// This choice has not yet been assigned
///
/// A value of `Unknown` is invalid after evaluation
Unknown = 0,
/// The operation always picks the left-hand input
Left = 1,
/// The operation always picks the right-hand input
Right = 2,
/// The operation may pick either input
Both = 3,
}
impl std::ops::BitOrAssign<Choice> for Choice {
fn bitor_assign(&mut self, other: Self) {
*self = match (*self as u8) | (other as u8) {
0 => Self::Unknown,
1 => Self::Left,
2 => Self::Right,
3 => Self::Both,
_ => unreachable!(),
}
}
}
impl std::ops::Not for Choice {
type Output = Choice;
fn not(self) -> Self {
match self {
Self::Unknown => Self::Both,
Self::Left => Self::Right,
Self::Right => Self::Left,
Self::Both => Self::Unknown,
}
}
}
impl std::ops::BitAndAssign<Choice> for Choice {
fn bitand_assign(&mut self, other: Self) {
*self = match (*self as u8) | ((!other as u8) & 0b11) {
0 => Self::Unknown,
1 => Self::Left,
2 => Self::Right,
3 => Self::Both,
_ => unreachable!(),
}
}
}
struct VChoice(u32x8);
impl VChoice {
pub const BOTH: Self = Self(u32x8::splat(Choice::Both as u32));
pub const LEFT: Self = Self(u32x8::splat(Choice::Left as u32));
pub const RIGHT: Self = Self(u32x8::splat(Choice::Right as u32));
pub const UNKNOWN: Self = Self(u32x8::splat(Choice::Unknown as u32));
}
impl std::ops::Not for VChoice {
type Output = Self;
fn not(self) -> Self {
VChoice(u32x8::splat(3) - self.0)
}
}
impl std::ops::BitAndAssign<VChoice> for VChoice {
fn bitand_assign(&mut self, other: Self) {
self.0 = (self.0) | ((!other.0) & u32x8::splat(3))
}
}
/// Stores a range, with conservative calculations to guarantee that it always
/// contains the actual value.
///
@@ -134,10 +225,22 @@ impl Interval {
Interval::new(lower, upper)
}
/// Cubes the interval
///
/// Note that this has tighter bounds than multiplication, because we know
/// that both sides of the multiplication are the same value.
#[inline]
pub fn cube(self) -> Self {
let has_nan = self.has_nan();
let lower = has_nan.select(VALUE_NAN, self.lower * self.lower * self.lower);
let upper = has_nan.select(VALUE_NAN, self.upper * self.upper * self.upper);
Interval::new(lower, upper)
}
/// Computes the sine of the interval
#[inline]
pub fn sin(self) -> Self {
let same_cycle = ((self.lower / VALUE_0) + VALUE_05)
let same_cycle = ((self.lower / VALUE_PI) + VALUE_05)
.floor()
.simd_eq(((self.upper / VALUE_PI) + VALUE_05).floor());
let up = (((self.upper / VALUE_PI) + VALUE_05).floor()) % (VALUE_2);
@@ -162,7 +265,7 @@ impl Interval {
/// Computes the cosine of the interval
#[inline]
pub fn cos(self) -> Self {
let same_cycle = (self.lower / VALUE_0)
let same_cycle = (self.lower / VALUE_PI)
.floor()
.simd_eq((self.upper / VALUE_PI).floor());
let up = ((self.upper / VALUE_PI).floor()) % (VALUE_2);
@@ -249,11 +352,9 @@ impl Interval {
/// Returns the `NAN` interval if the input contains zero
#[inline]
pub fn ln(self) -> Self {
if self.lower <= 0.0 {
f32::NAN.into()
} else {
Interval::new(self.lower.ln(), self.upper.ln())
}
let lower = (self.has_nan()).select(VALUE_NAN, self.lower.ln());
let upper = (self.has_nan()).select(VALUE_NAN, self.upper.ln());
Interval::new(lower, upper)
}
/// Calculates the square root of the interval
@@ -261,11 +362,9 @@ impl Interval {
/// If the interval contains values below 0, returns a `NAN` interval.
#[inline]
pub fn sqrt(self) -> Self {
if self.lower < 0.0 {
f32::NAN.into()
} else {
Interval::new(self.lower.sqrt(), self.upper.sqrt())
}
let lower = (self.lower.simd_lt(VALUE_0)).select(VALUE_NAN, self.lower.sqrt());
let upper = (self.lower.simd_lt(VALUE_0)).select(VALUE_NAN, self.upper.sqrt());
Interval::new(lower, upper)
}
/// Calculates the reciprocal of the interval
@@ -273,110 +372,149 @@ impl Interval {
/// If the interval includes 0, returns the `NAN` interval
#[inline]
pub fn recip(self) -> Self {
if self.lower > 0.0 || self.upper < 0.0 {
Interval::new(1.0 / self.upper, 1.0 / self.lower)
} else {
f32::NAN.into()
}
let lower = (self.lower.simd_le(VALUE_0) & self.upper.simd_ge(VALUE_0))
.select(VALUE_NAN, self.upper.recip());
let upper = (self.lower.simd_le(VALUE_0) & self.upper.simd_ge(VALUE_0))
.select(VALUE_NAN, self.lower.recip());
Interval::new(lower, upper)
}
/// Calculates the minimum of two intervals
///
/// Returns both the result and a [`Choice`] indicating whether one side is
/// Returns both the result and a [`VChoice`] indicating whether one side is
/// always less than the other.
///
/// If either side is `NAN`, returns the `NAN` interval and `Choice::Both`.
/// If either side is `NAN`, returns the `NAN` interval and `VChoice::Both`.
#[inline]
pub fn min_choice(self, rhs: Self) -> (Self, Choice) {
if self.has_nan() || rhs.has_nan() {
return (f32::NAN.into(), Choice::Both);
}
let choice = if self.upper < rhs.lower {
Choice::Left
} else if rhs.upper < self.lower {
Choice::Right
} else {
Choice::Both
};
pub fn min_choice(self, rhs: Self) -> (Self, VChoice) {
let has_nan = self.has_nan() | rhs.has_nan();
let choice = has_nan.select(
VChoice::BOTH.0,
self.upper.simd_lt(rhs.lower).select(
VChoice::LEFT.0,
rhs.upper
.simd_lt(self.lower)
.select(VChoice::RIGHT.0, VChoice::BOTH.0),
),
);
(
Interval::new(self.lower.min(rhs.lower), self.upper.min(rhs.upper)),
choice,
Interval::new(
has_nan.select(VALUE_NAN, self.lower.simd_min(rhs.lower)),
has_nan.select(VALUE_NAN, self.upper.simd_min(rhs.upper)),
),
VChoice(choice),
)
}
/// Calculates the maximum of two intervals
///
/// Returns both the result and a [`Choice`] indicating whether one side is
/// Returns both the result and a [`VChoice`] indicating whether one side is
/// always greater than the other.
///
/// If either side is `NAN`, returns the `NAN` interval and `Choice::Both`.
/// If either side is `NAN`, returns the `NAN` interval and `VChoice::Both`.
#[inline]
pub fn max_choice(self, rhs: Self) -> (Self, Choice) {
if self.has_nan() || rhs.has_nan() {
return (f32::NAN.into(), Choice::Both);
}
let choice = if self.lower > rhs.upper {
Choice::Left
} else if rhs.lower > self.upper {
Choice::Right
} else {
Choice::Both
};
pub fn max_choice(self, rhs: Self) -> (Self, VChoice) {
let has_nan = self.has_nan() | rhs.has_nan();
let choice = has_nan.select(
VChoice::BOTH.0,
self.lower.simd_gt(rhs.upper).select(
VChoice::LEFT.0,
rhs.lower
.simd_gt(self.upper)
.select(VChoice::RIGHT.0, VChoice::BOTH.0),
),
);
(
Interval::new(self.lower.max(rhs.lower), self.upper.max(rhs.upper)),
choice,
Interval::new(
has_nan.select(VALUE_NAN, self.lower.simd_min(rhs.lower)),
has_nan.select(VALUE_NAN, self.upper.simd_min(rhs.upper)),
),
VChoice(choice),
)
}
/// Calculates the short-circuiting `AND` of two intervals
///
/// Returns both the result and a [`Choice`] indicating whether one side is
/// Returns both the result and a [`VChoice`] indicating whether one side is
/// always selected. An unambiguous 0 in `self` selects itself; an
/// unambiguous 1 selects the opposite branch.
#[inline]
pub fn and_choice(self, rhs: Self) -> (Self, Choice) {
if self.has_nan() || rhs.has_nan() {
(f32::NAN.into(), Choice::Both)
} else if self.lower == 0.0 && self.upper == 0.0 {
(0.0.into(), Choice::Left)
} else if !self.contains(0.0) {
(rhs, Choice::Right)
} else {
// The output will either be the RHS or zero, so extend the interval
// to include zero in it.
(
Interval::new(rhs.lower.min(0.0), rhs.upper.max(0.0)),
Choice::Both,
)
}
pub fn and_choice(self, rhs: Self) -> (Self, VChoice) {
let has_nan = self.has_nan() | rhs.has_nan();
let choice = has_nan.select(
VChoice::BOTH.0,
(self.lower.simd_eq(VALUE_0) & self.upper.simd_eq(VALUE_0)).select(
VChoice::LEFT.0,
self.contains(VALUE_0)
.select(VChoice::BOTH.0, VChoice::RIGHT.0),
),
);
(
Interval::new(
has_nan.select(
VALUE_NAN,
(self.lower.simd_eq(VALUE_0) & self.upper.simd_eq(VALUE_0)).select(
VALUE_0,
self.contains(VALUE_0)
.select(rhs.lower.simd_min(VALUE_0), rhs.lower),
),
),
has_nan.select(
VALUE_NAN,
(self.lower.simd_eq(VALUE_0) & self.upper.simd_eq(VALUE_0)).select(
VALUE_0,
self.contains(VALUE_0)
.select(rhs.upper.simd_max(VALUE_0), rhs.upper),
),
),
),
VChoice(choice),
)
}
/// Calculates the short-circuiting `OR` of two intervals
///
/// Returns both the result and a [`Choice`] indicating whether one side is
/// Returns both the result and a [`VChoice`] indicating whether one side is
/// always selected. An unambiguous 0 in `self` selects the opposite
/// branch; an unambiguous 1 selects itself.
#[inline]
pub fn or_choice(self, rhs: Self) -> (Self, Choice) {
if self.has_nan() || rhs.has_nan() {
(f32::NAN.into(), Choice::Both)
} else if !self.contains(0.0) {
(self, Choice::Left)
} else if self.lower == 0.0 && self.upper == 0.0 {
(rhs, Choice::Right)
} else {
// The output could be anywhere in either interval
(
Interval::new(self.lower.min(rhs.lower), self.upper.max(rhs.upper)),
Choice::Both,
)
}
pub fn or_choice(self, rhs: Self) -> (Self, VChoice) {
let has_nan = self.has_nan() | rhs.has_nan();
let choice = has_nan.select(
VChoice::BOTH.0,
self.contains(VALUE_0).select(
(self.lower.simd_eq(VALUE_0) & self.upper.simd_eq(VALUE_0))
.select(VChoice::RIGHT.0, VChoice::BOTH.0),
VChoice::LEFT.0,
),
);
(
Interval::new(
has_nan.select(
VALUE_NAN,
self.contains(VALUE_0).select(
(self.lower.simd_eq(VALUE_0) & self.upper.simd_eq(VALUE_0))
.select(rhs.lower, rhs.lower.simd_min(self.lower)),
self.lower,
),
),
has_nan.select(
VALUE_NAN,
self.contains(VALUE_0).select(
(self.lower.simd_eq(VALUE_0) & self.upper.simd_eq(VALUE_0))
.select(rhs.upper, rhs.upper.simd_max(self.upper)),
self.upper,
),
),
),
VChoice(choice),
)
}
/// Returns the midpoint of the interval
#[inline]
pub fn midpoint(self) -> f32 {
(self.lower + self.upper) / 2.0
pub fn midpoint(self) -> Value {
(self.lower + self.upper) / VALUE_2
}
/// Splits the interval at the midpoint
@@ -407,8 +545,8 @@ impl Interval {
/// assert_eq!(a.lerp(2.0), 4.0);
/// ```
#[inline]
pub fn lerp(self, frac: f32) -> f32 {
self.lower * (1.0 - frac) + self.upper * frac
pub fn lerp(self, frac: Value) -> Value {
self.lower * (VALUE_1 - frac) + self.upper * frac
}
/// Calculates the width of the interval
@@ -421,7 +559,7 @@ impl Interval {
/// assert_eq!(b.width(), 3.0);
/// ```
#[inline]
pub fn width(self) -> f32 {
pub fn width(self) -> Value {
self.upper - self.lower
}
@@ -430,8 +568,8 @@ impl Interval {
pub(crate) fn compare_eq(&self, other: Self) {
let d = (self.lower - other.lower)
.abs()
.max((self.upper - other.upper).abs());
if d >= 1e-6 {
.simd_max((self.upper - other.upper).abs());
if d.simd_ge(Value::splat(1e-6)).any() {
panic!("lhs != rhs ({self:?} != {other:?})");
}
}
@@ -440,22 +578,22 @@ impl Interval {
#[inline]
pub fn rem_euclid(&self, other: Interval) -> Self {
// TODO optimize this more?
if self.has_nan() || other.has_nan() || other.contains(0.0) {
f32::NAN.into()
} else if other.lower == other.upper && other.lower > 0.0 {
let a = self.lower / other.lower;
let b = self.upper / other.lower;
if a != a.floor() && a.floor() == b.floor() {
Interval::new(
self.lower.rem_euclid(other.lower),
self.upper.rem_euclid(other.lower),
)
} else {
Interval::new(0.0, other.abs().upper())
}
} else {
Interval::new(0.0, other.abs().upper())
}
let has_nan = self.has_nan() | other.has_nan() | other.contains(VALUE_0);
let other_constant = other.lower.simd_eq(other.upper) & other.lower.simd_gt(VALUE_0);
let a = self.lower / other.lower;
let b = self.upper / other.lower;
let floors = a.simd_ne(a.floor()) & a.floor().simd_eq(b.floor());
let lower = has_nan.select(
VALUE_NAN,
(other_constant & floors).select(self.lower % other.lower, VALUE_0),
);
let upper = has_nan.select(
VALUE_NAN,
(other_constant & floors).select(self.upper % other.lower, other.upper.abs()),
);
Interval::new(lower, upper)
}
/// Largest value that is less-than-or-equal to this value
@@ -479,31 +617,31 @@ impl Interval {
/// Four-quadrant arctangent
#[inline]
pub fn atan2(self, x: Self) -> Self {
if self.has_nan() || x.has_nan() {
f32::NAN.into()
} else {
// TODO optimize this further
Interval::new(-std::f32::consts::PI, std::f32::consts::PI)
}
let has_nan = self.has_nan() | x.has_nan();
// TODO optimize this further
Interval::new(
has_nan.select(VALUE_NAN, -VALUE_PI),
has_nan.select(VALUE_NAN, VALUE_PI),
)
}
}
impl std::fmt::Display for Interval {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "({}, {})", self.lower, self.upper)
write!(f, "({:?}, {:?})", self.lower, self.upper)
}
}
impl From<[f32; 2]> for Interval {
impl From<[Value; 2]> for Interval {
#[inline]
fn from(i: [f32; 2]) -> Interval {
fn from(i: [Value; 2]) -> Interval {
Interval::new(i[0], i[1])
}
}
impl From<f32> for Interval {
impl From<Value> for Interval {
#[inline]
fn from(f: f32) -> Self {
fn from(f: Value) -> Self {
Interval::new(f, f)
}
}
@@ -522,10 +660,8 @@ impl std::ops::Mul<Interval> for Interval {
#[inline]
fn mul(self, rhs: Self) -> Self {
if self.has_nan() || rhs.has_nan() {
return f32::NAN.into();
}
let mut out = [0.0; 4];
let has_nan = self.has_nan() | rhs.has_nan();
let mut out = [VALUE_0; 4];
let mut k = 0;
for i in [self.lower, self.upper] {
for j in [rhs.lower, rhs.upper] {
@@ -536,25 +672,27 @@ impl std::ops::Mul<Interval> for Interval {
let mut lower = out[0];
let mut upper = out[0];
for &v in &out[1..] {
lower = lower.min(v);
upper = upper.max(v);
lower = lower.simd_min(v);
upper = upper.simd_max(v);
}
Interval::new(lower, upper)
Interval::new(
has_nan.select(VALUE_NAN, lower),
has_nan.select(VALUE_NAN, upper),
)
}
}
impl std::ops::Mul<f32> for Interval {
impl std::ops::Mul<Value> for Interval {
type Output = Self;
#[inline]
fn mul(self, rhs: f32) -> Self {
if self.has_nan() || rhs.is_nan() {
f32::NAN.into()
} else if rhs < 0.0 {
Interval::new(self.upper * rhs, self.lower * rhs)
} else {
Interval::new(self.lower * rhs, self.upper * rhs)
}
fn mul(self, rhs: Value) -> Self {
let has_nan = self.has_nan() | rhs.is_nan();
let rlt = rhs.simd_lt(VALUE_0);
Interval::new(
has_nan.select(VALUE_NAN, rlt.select(self.upper * rhs, self.lower * rhs)),
has_nan.select(VALUE_NAN, rlt.select(self.lower * rhs, self.upper * rhs)),
)
}
}
@@ -563,28 +701,25 @@ impl std::ops::Div<Interval> for Interval {
#[inline]
fn div(self, rhs: Self) -> Self {
if self.has_nan() {
return f32::NAN.into();
}
if rhs.lower > 0.0 || rhs.upper < 0.0 {
let mut out = [0.0; 4];
let mut k = 0;
for i in [self.lower, self.upper] {
for j in [rhs.lower, rhs.upper] {
out[k] = i / j;
k += 1;
}
let has_nan = self.has_nan() | (rhs.lower.simd_lt(VALUE_0) & rhs.upper.simd_gt(VALUE_0));
let mut out = [VALUE_0; 4];
let mut k = 0;
for i in [self.lower, self.upper] {
for j in [rhs.lower, rhs.upper] {
out[k] = i / j;
k += 1;
}
let mut lower = out[0];
let mut upper = out[0];
for &v in &out[1..] {
lower = lower.min(v);
upper = upper.max(v);
}
Interval::new(lower, upper)
} else {
f32::NAN.into()
}
let mut lower = out[0];
let mut upper = out[0];
for &v in &out[1..] {
lower = lower.simd_min(v);
upper = upper.simd_max(v);
}
Interval::new(
has_nan.select(VALUE_NAN, lower),
has_nan.select(VALUE_NAN, upper),
)
}
}
@@ -612,10 +747,10 @@ mod test {
#[test]
fn test_interval() {
let a = Interval::new(0.0, 1.0);
let b = Interval::new(0.5, 1.5);
let a = Interval::new(Value::splat(0.0), Value::splat(1.0));
let b = Interval::new(Value::splat(0.5), Value::splat(1.5));
let (v, c) = a.min_choice(b);
assert_eq!(v, [0.0, 1.0].into());
assert_eq!(c, Choice::Both);
assert_eq!(v, [Value::splat(0.0), Value::splat(1.0)].into());
assert_eq!(c.0, VChoice::BOTH.0);
}
}