--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, StdFloat,
cmp::{SimdPartialEq, SimdPartialOrd}, cmp::{SimdPartialEq, SimdPartialOrd},
num::SimdFloat, num::SimdFloat,
u32x8,
}; };
use crate::interpreter::{ use crate::interpreter::{
Mask, VALUE_0, VALUE_1, VALUE_2, VALUE_05, VALUE_M1, VALUE_NAN, VALUE_PI, Value, 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 /// Stores a range, with conservative calculations to guarantee that it always
/// contains the actual value. /// contains the actual value.
/// ///
@@ -134,10 +225,22 @@ impl Interval {
Interval::new(lower, upper) 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 /// Computes the sine of the interval
#[inline] #[inline]
pub fn sin(self) -> Self { pub fn sin(self) -> Self {
let same_cycle = ((self.lower / VALUE_0) + VALUE_05) let same_cycle = ((self.lower / VALUE_PI) + VALUE_05)
.floor() .floor()
.simd_eq(((self.upper / 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); let up = (((self.upper / VALUE_PI) + VALUE_05).floor()) % (VALUE_2);
@@ -162,7 +265,7 @@ impl Interval {
/// Computes the cosine of the interval /// Computes the cosine of the interval
#[inline] #[inline]
pub fn cos(self) -> Self { pub fn cos(self) -> Self {
let same_cycle = (self.lower / VALUE_0) let same_cycle = (self.lower / VALUE_PI)
.floor() .floor()
.simd_eq((self.upper / VALUE_PI).floor()); .simd_eq((self.upper / VALUE_PI).floor());
let up = ((self.upper / VALUE_PI).floor()) % (VALUE_2); let up = ((self.upper / VALUE_PI).floor()) % (VALUE_2);
@@ -249,11 +352,9 @@ impl Interval {
/// Returns the `NAN` interval if the input contains zero /// Returns the `NAN` interval if the input contains zero
#[inline] #[inline]
pub fn ln(self) -> Self { pub fn ln(self) -> Self {
if self.lower <= 0.0 { let lower = (self.has_nan()).select(VALUE_NAN, self.lower.ln());
f32::NAN.into() let upper = (self.has_nan()).select(VALUE_NAN, self.upper.ln());
} else { Interval::new(lower, upper)
Interval::new(self.lower.ln(), self.upper.ln())
}
} }
/// Calculates the square root of the interval /// Calculates the square root of the interval
@@ -261,11 +362,9 @@ impl Interval {
/// If the interval contains values below 0, returns a `NAN` interval. /// If the interval contains values below 0, returns a `NAN` interval.
#[inline] #[inline]
pub fn sqrt(self) -> Self { pub fn sqrt(self) -> Self {
if self.lower < 0.0 { let lower = (self.lower.simd_lt(VALUE_0)).select(VALUE_NAN, self.lower.sqrt());
f32::NAN.into() let upper = (self.lower.simd_lt(VALUE_0)).select(VALUE_NAN, self.upper.sqrt());
} else { Interval::new(lower, upper)
Interval::new(self.lower.sqrt(), self.upper.sqrt())
}
} }
/// Calculates the reciprocal of the interval /// Calculates the reciprocal of the interval
@@ -273,110 +372,149 @@ impl Interval {
/// If the interval includes 0, returns the `NAN` interval /// If the interval includes 0, returns the `NAN` interval
#[inline] #[inline]
pub fn recip(self) -> Self { pub fn recip(self) -> Self {
if self.lower > 0.0 || self.upper < 0.0 { let lower = (self.lower.simd_le(VALUE_0) & self.upper.simd_ge(VALUE_0))
Interval::new(1.0 / self.upper, 1.0 / self.lower) .select(VALUE_NAN, self.upper.recip());
} else { let upper = (self.lower.simd_le(VALUE_0) & self.upper.simd_ge(VALUE_0))
f32::NAN.into() .select(VALUE_NAN, self.lower.recip());
} Interval::new(lower, upper)
} }
/// Calculates the minimum of two intervals /// 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. /// 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] #[inline]
pub fn min_choice(self, rhs: Self) -> (Self, Choice) { pub fn min_choice(self, rhs: Self) -> (Self, VChoice) {
if self.has_nan() || rhs.has_nan() { let has_nan = self.has_nan() | rhs.has_nan();
return (f32::NAN.into(), Choice::Both); let choice = has_nan.select(
} VChoice::BOTH.0,
let choice = if self.upper < rhs.lower { self.upper.simd_lt(rhs.lower).select(
Choice::Left VChoice::LEFT.0,
} else if rhs.upper < self.lower { rhs.upper
Choice::Right .simd_lt(self.lower)
} else { .select(VChoice::RIGHT.0, VChoice::BOTH.0),
Choice::Both ),
}; );
( (
Interval::new(self.lower.min(rhs.lower), self.upper.min(rhs.upper)), Interval::new(
choice, 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 /// 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. /// 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] #[inline]
pub fn max_choice(self, rhs: Self) -> (Self, Choice) { pub fn max_choice(self, rhs: Self) -> (Self, VChoice) {
if self.has_nan() || rhs.has_nan() { let has_nan = self.has_nan() | rhs.has_nan();
return (f32::NAN.into(), Choice::Both); let choice = has_nan.select(
} VChoice::BOTH.0,
let choice = if self.lower > rhs.upper { self.lower.simd_gt(rhs.upper).select(
Choice::Left VChoice::LEFT.0,
} else if rhs.lower > self.upper { rhs.lower
Choice::Right .simd_gt(self.upper)
} else { .select(VChoice::RIGHT.0, VChoice::BOTH.0),
Choice::Both ),
}; );
( (
Interval::new(self.lower.max(rhs.lower), self.upper.max(rhs.upper)), Interval::new(
choice, 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 /// 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 /// always selected. An unambiguous 0 in `self` selects itself; an
/// unambiguous 1 selects the opposite branch. /// unambiguous 1 selects the opposite branch.
#[inline] #[inline]
pub fn and_choice(self, rhs: Self) -> (Self, Choice) { pub fn and_choice(self, rhs: Self) -> (Self, VChoice) {
if self.has_nan() || rhs.has_nan() { let has_nan = self.has_nan() | rhs.has_nan();
(f32::NAN.into(), Choice::Both) let choice = has_nan.select(
} else if self.lower == 0.0 && self.upper == 0.0 { VChoice::BOTH.0,
(0.0.into(), Choice::Left) (self.lower.simd_eq(VALUE_0) & self.upper.simd_eq(VALUE_0)).select(
} else if !self.contains(0.0) { VChoice::LEFT.0,
(rhs, Choice::Right) self.contains(VALUE_0)
} else { .select(VChoice::BOTH.0, VChoice::RIGHT.0),
// 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)), Interval::new(
Choice::Both, 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 /// 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 /// always selected. An unambiguous 0 in `self` selects the opposite
/// branch; an unambiguous 1 selects itself. /// branch; an unambiguous 1 selects itself.
#[inline] #[inline]
pub fn or_choice(self, rhs: Self) -> (Self, Choice) { pub fn or_choice(self, rhs: Self) -> (Self, VChoice) {
if self.has_nan() || rhs.has_nan() { let has_nan = self.has_nan() | rhs.has_nan();
(f32::NAN.into(), Choice::Both) let choice = has_nan.select(
} else if !self.contains(0.0) { VChoice::BOTH.0,
(self, Choice::Left) self.contains(VALUE_0).select(
} else if self.lower == 0.0 && self.upper == 0.0 { (self.lower.simd_eq(VALUE_0) & self.upper.simd_eq(VALUE_0))
(rhs, Choice::Right) .select(VChoice::RIGHT.0, VChoice::BOTH.0),
} else { VChoice::LEFT.0,
// The output could be anywhere in either interval ),
( );
Interval::new(self.lower.min(rhs.lower), self.upper.max(rhs.upper)), (
Choice::Both, 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 /// Returns the midpoint of the interval
#[inline] #[inline]
pub fn midpoint(self) -> f32 { pub fn midpoint(self) -> Value {
(self.lower + self.upper) / 2.0 (self.lower + self.upper) / VALUE_2
} }
/// Splits the interval at the midpoint /// Splits the interval at the midpoint
@@ -407,8 +545,8 @@ impl Interval {
/// assert_eq!(a.lerp(2.0), 4.0); /// assert_eq!(a.lerp(2.0), 4.0);
/// ``` /// ```
#[inline] #[inline]
pub fn lerp(self, frac: f32) -> f32 { pub fn lerp(self, frac: Value) -> Value {
self.lower * (1.0 - frac) + self.upper * frac self.lower * (VALUE_1 - frac) + self.upper * frac
} }
/// Calculates the width of the interval /// Calculates the width of the interval
@@ -421,7 +559,7 @@ impl Interval {
/// assert_eq!(b.width(), 3.0); /// assert_eq!(b.width(), 3.0);
/// ``` /// ```
#[inline] #[inline]
pub fn width(self) -> f32 { pub fn width(self) -> Value {
self.upper - self.lower self.upper - self.lower
} }
@@ -430,8 +568,8 @@ impl Interval {
pub(crate) fn compare_eq(&self, other: Self) { pub(crate) fn compare_eq(&self, other: Self) {
let d = (self.lower - other.lower) let d = (self.lower - other.lower)
.abs() .abs()
.max((self.upper - other.upper).abs()); .simd_max((self.upper - other.upper).abs());
if d >= 1e-6 { if d.simd_ge(Value::splat(1e-6)).any() {
panic!("lhs != rhs ({self:?} != {other:?})"); panic!("lhs != rhs ({self:?} != {other:?})");
} }
} }
@@ -440,22 +578,22 @@ impl Interval {
#[inline] #[inline]
pub fn rem_euclid(&self, other: Interval) -> Self { pub fn rem_euclid(&self, other: Interval) -> Self {
// TODO optimize this more? // TODO optimize this more?
if self.has_nan() || other.has_nan() || other.contains(0.0) { let has_nan = self.has_nan() | other.has_nan() | other.contains(VALUE_0);
f32::NAN.into() let other_constant = other.lower.simd_eq(other.upper) & other.lower.simd_gt(VALUE_0);
} else if other.lower == other.upper && other.lower > 0.0 { let a = self.lower / other.lower;
let a = self.lower / other.lower; let b = self.upper / other.lower;
let b = self.upper / other.lower; let floors = a.simd_ne(a.floor()) & a.floor().simd_eq(b.floor());
if a != a.floor() && a.floor() == b.floor() {
Interval::new( let lower = has_nan.select(
self.lower.rem_euclid(other.lower), VALUE_NAN,
self.upper.rem_euclid(other.lower), (other_constant & floors).select(self.lower % other.lower, VALUE_0),
) );
} else { let upper = has_nan.select(
Interval::new(0.0, other.abs().upper()) VALUE_NAN,
} (other_constant & floors).select(self.upper % other.lower, other.upper.abs()),
} else { );
Interval::new(0.0, other.abs().upper())
} Interval::new(lower, upper)
} }
/// Largest value that is less-than-or-equal to this value /// Largest value that is less-than-or-equal to this value
@@ -479,31 +617,31 @@ impl Interval {
/// Four-quadrant arctangent /// Four-quadrant arctangent
#[inline] #[inline]
pub fn atan2(self, x: Self) -> Self { pub fn atan2(self, x: Self) -> Self {
if self.has_nan() || x.has_nan() { let has_nan = self.has_nan() | x.has_nan();
f32::NAN.into() // TODO optimize this further
} else { Interval::new(
// TODO optimize this further has_nan.select(VALUE_NAN, -VALUE_PI),
Interval::new(-std::f32::consts::PI, std::f32::consts::PI) has_nan.select(VALUE_NAN, VALUE_PI),
} )
} }
} }
impl std::fmt::Display for Interval { impl std::fmt::Display for Interval {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { 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] #[inline]
fn from(i: [f32; 2]) -> Interval { fn from(i: [Value; 2]) -> Interval {
Interval::new(i[0], i[1]) Interval::new(i[0], i[1])
} }
} }
impl From<f32> for Interval { impl From<Value> for Interval {
#[inline] #[inline]
fn from(f: f32) -> Self { fn from(f: Value) -> Self {
Interval::new(f, f) Interval::new(f, f)
} }
} }
@@ -522,10 +660,8 @@ impl std::ops::Mul<Interval> for Interval {
#[inline] #[inline]
fn mul(self, rhs: Self) -> Self { fn mul(self, rhs: Self) -> Self {
if self.has_nan() || rhs.has_nan() { let has_nan = self.has_nan() | rhs.has_nan();
return f32::NAN.into(); let mut out = [VALUE_0; 4];
}
let mut out = [0.0; 4];
let mut k = 0; let mut k = 0;
for i in [self.lower, self.upper] { for i in [self.lower, self.upper] {
for j in [rhs.lower, rhs.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 lower = out[0];
let mut upper = out[0]; let mut upper = out[0];
for &v in &out[1..] { for &v in &out[1..] {
lower = lower.min(v); lower = lower.simd_min(v);
upper = upper.max(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; type Output = Self;
#[inline] #[inline]
fn mul(self, rhs: f32) -> Self { fn mul(self, rhs: Value) -> Self {
if self.has_nan() || rhs.is_nan() { let has_nan = self.has_nan() | rhs.is_nan();
f32::NAN.into() let rlt = rhs.simd_lt(VALUE_0);
} else if rhs < 0.0 { Interval::new(
Interval::new(self.upper * rhs, self.lower * rhs) has_nan.select(VALUE_NAN, rlt.select(self.upper * rhs, self.lower * rhs)),
} else { has_nan.select(VALUE_NAN, rlt.select(self.lower * rhs, self.upper * rhs)),
Interval::new(self.lower * rhs, self.upper * rhs) )
}
} }
} }
@@ -563,28 +701,25 @@ impl std::ops::Div<Interval> for Interval {
#[inline] #[inline]
fn div(self, rhs: Self) -> Self { fn div(self, rhs: Self) -> Self {
if self.has_nan() { let has_nan = self.has_nan() | (rhs.lower.simd_lt(VALUE_0) & rhs.upper.simd_gt(VALUE_0));
return f32::NAN.into(); let mut out = [VALUE_0; 4];
} let mut k = 0;
if rhs.lower > 0.0 || rhs.upper < 0.0 { for i in [self.lower, self.upper] {
let mut out = [0.0; 4]; for j in [rhs.lower, rhs.upper] {
let mut k = 0; out[k] = i / j;
for i in [self.lower, self.upper] { k += 1;
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] #[test]
fn test_interval() { fn test_interval() {
let a = Interval::new(0.0, 1.0); let a = Interval::new(Value::splat(0.0), Value::splat(1.0));
let b = Interval::new(0.5, 1.5); let b = Interval::new(Value::splat(0.5), Value::splat(1.5));
let (v, c) = a.min_choice(b); let (v, c) = a.min_choice(b);
assert_eq!(v, [0.0, 1.0].into()); assert_eq!(v, [Value::splat(0.0), Value::splat(1.0)].into());
assert_eq!(c, Choice::Both); assert_eq!(c.0, VChoice::BOTH.0);
} }
} }