Move to workspace

This commit is contained in:
2025-12-20 22:04:47 +00:00
parent 8fcd512c36
commit b6b0c2114f
41 changed files with 4540 additions and 3402 deletions
+210
View File
@@ -0,0 +1,210 @@
#[repr(u8)]
#[allow(non_snake_case)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum InstructionSet {
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns the input. Useful for copying registers.
OPCopy = (1),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Adds a vector to a vector component-wise.
OPAdd = (2),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Subtracts a vector from a vector component-wise.
OPSub = (3),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Multiplies a vector and a vector component-wise.
OPMul = (4),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Divides a vector by a vector component-wise.
OPDiv = (5),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Calculates a vector Atan2 a vector component-wise.
OPAtan2 = (6),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Calculates the minimum of a vector and a vector component-wise.
OPMin = (7),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Calculates the maximum of a vector and a vector component-wise.
OPMax = (8),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Threeway comparison operator.
OPCompare = (9),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Calculates a vector modulo a vector component-wise.
OPMod = (10),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// If both arguments are non-zero, returns the right-hand argument.
/// Otherwise, returns zero.
OPAnd = (11),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// If the left-hand argument is non-zero, it is returned. Otherwise, the
/// right-hand argument is returned.
OPOr = (12),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns the negation of all components of a vector.
OPNegate = (13),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns the absolute value of all components of a vector.
OPAbs = (14),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns 1 over all components of a vector.
OPRecip = (15),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns the square root of all components of a vector.
OPSqrt = (16),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns the square of all components of a vector.
OPSquare = (17),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns the floor of all components of a vector.
OPFloor = (18),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns the ceiling of all components of a vector.
OPCeil = (19),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns all components of a vector rounded to the nearest integer, 0.5
/// away from zero.
OPRound = (20),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns the sine of all components of a vector.
OPSin = (21),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns the cosine of all components of a vector.
OPCos = (22),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns the tangent of all components of a vector.
OPTan = (23),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns the arc sine of all components of a vector.
OPAsin = (24),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns the arc cosine of all components of a vector.
OPAcos = (25),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns the arc tangent of all components of a vector.
OPAtan = (26),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns e raised to all components of a vector.
OPExp = (27),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns the natural logarithm of all components of a vector.
OPLog = (28),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// The output is 1 if the argument is 0, and 0 otherwise.
OPNot = (29),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns the fractional part of all components of a vector.
OPFract = (30),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns the cube of all components of a vector.
OPCube = (31),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns the smooth minimum between a vector and a vector, varied by a
/// vector.
OPSmoothMin = (32),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns the smooth maximum between a vector and a vector, varied by a
/// vector.
OPSmoothMax = (33),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Clamps a vector between a vector and a vector.
OPClamp = (34),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Mixes between a vector and a vector, varied by a vector.
OPMix = (35),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Calculates a vector multiplied by a vector, then adds a vector.
OPFMA = (36),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns the length (magnitude) of a vector.
OPLength = (37),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns the dot product of two vectors.
OPDot = (38),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns the length (magnitude) of the vector between two vectors.
OPDistance = (39),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// No operation.
OPNop = ((3 * 64) + 63),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Stops execution of the tape and returns a single value.
OPReturn = ((2 * 64) + 63),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns the current position being sampled.
OPPosition = ((1 * 64) + 63),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Calculates the minimum of two Vec1s, and also carries over the relevant
/// material metadata.
OPMinMaterial = ((0 * 64) + 63),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Calculates the maximum of two Vec1s, and also carries over the relevant
/// material metadata.
OPMaxMaterial = ((3 * 64) + 62),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns the smooth minimum between a Vec1 and a Vec1, varied by a Vec1,
/// and also carries over the relevant material metadata.
OPSmoothMinMaterial = ((2 * 64) + 62),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns the smooth maximum between a Vec1 and a Vec1, varied by a Vec1,
/// and also carries over the relevant material metadata.
OPSmoothMaxMaterial = ((1 * 64) + 62),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns the distance to a sphere.
OPSDFSphere = ((3 * 64) + 61),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns the distance to a box.
OPSDFBox = ((2 * 64) + 61),
#[allow(non_snake_case)]
#[allow(dead_code)]
/// Returns the distance to a torus.
OPSDFTorus = ((1 * 64) + 61),
}
+28
View File
@@ -0,0 +1,28 @@
# Interpreter redesign
Ground up redesign of the interpreter
Maximum mesh shaders per SM is 128. For 65,536 vgprs, each mesh shader can use 512.
Instead of a stack, use SSA and then limited registers. This is a bit more compile friendly.
There are 16 available registers. Register 0 is always 0, and register 15 is the next item in the const tape.
This leaves 14 usable 32 bit float registers.
Instruction format:
8 bits
If the lowest 6 bits are below 47, the top two bits are length of the VecX type, where 0 means Vec1 and 3 means Vec4.
Otherwise, all 8 bits encode an instruction.
Each following 8 bits encode two 4 bit registers that are registers. All inputs are listed, then all outputs. It is padded to 8 bits with zeros.
Each SDF can have up to 8 materials associated with it. The relative weight of each material is tracked through the interpreter as an
array of 8 floats, and at the end they are evaluated once.
Steps:
1. 8x8x1 task shaders per mesh (64)
2. each task shader works out 4x4x2 blocks (32) where 2 is depth
3. each task shader spawns 2x2x32 (128) mesh shaders per block.
4. each subgroup outputs 4 verts and 2 triangles, for a total 16/8 per workgroup.
For a mesh that takes up 1024x1024 pixel on screen, each quad takes up 16x16 pixels.
+419
View File
@@ -0,0 +1,419 @@
use std::simd::{StdFloat, cmp::SimdPartialOrd, num::SimdFloat};
use crate::{
interpreters::VALUE_0,
ssa::{SSAInput, SSAInstruction, SSAOpcode, SSATape},
types::Interval,
};
#[derive(Clone, Debug)]
pub struct IntervalInterpreter<'csg> {
value_map: Vec<Interval>,
csg: &'csg SSATape,
}
impl IntervalInterpreter<'_> {
fn load(&self, consider: SSAInput) -> Interval {
match consider {
SSAInput::Register(r) => self.value_map[r as usize],
SSAInput::Constant(c) => Interval::splat(c),
}
}
fn store(&mut self, location: u32, value: Interval) {
self.value_map[location as usize] = value;
}
}
impl<'csg> IntervalInterpreter<'csg> {
pub fn new(csg: &'csg SSATape) -> Self {
IntervalInterpreter {
value_map: vec![Interval::ZERO; csg.last_output as usize],
csg,
}
}
fn clear_stacks(&mut self) {
self.value_map = vec![Interval::ZERO; self.csg.last_output as usize];
}
fn param_one(&mut self, instruction: &SSAInstruction, func: impl Fn(Interval) -> Interval) {
for i in 0..instruction.opcode.size as usize {
let val_a = self.load(instruction.inputs[i]);
self.store(instruction.outputs[i], func(val_a));
}
}
fn param_two(
&mut self,
instruction: &SSAInstruction,
func: impl Fn(Interval, Interval) -> Interval,
) {
for i in 0..instruction.opcode.size as usize {
let val_a = self.load(instruction.inputs[i]);
let val_b = self.load(instruction.inputs[i + instruction.opcode.size as usize]);
self.store(instruction.outputs[i], func(val_a, val_b));
}
}
fn param_three(
&mut self,
instruction: &SSAInstruction,
func: impl Fn(Interval, Interval, Interval) -> Interval,
) {
for i in 0..instruction.opcode.size as usize {
let val_a = self.load(instruction.inputs[i]);
let val_b = self.load(instruction.inputs[i + instruction.opcode.size as usize]);
let val_c = self.load(instruction.inputs[i + (instruction.opcode.size * 2) as usize]);
self.store(instruction.outputs[i], func(val_a, val_b, val_c));
}
}
fn param_four(
&mut self,
instruction: &SSAInstruction,
func: impl Fn(Interval, Interval, Interval, Interval) -> Interval,
) {
for i in 0..instruction.opcode.size as usize {
let val_a = self.load(instruction.inputs[i]);
let val_b = self.load(instruction.inputs[i + instruction.opcode.size as usize]);
let val_c = self.load(instruction.inputs[i + (instruction.opcode.size * 2) as usize]);
let val_d = self.load(instruction.inputs[i + (instruction.opcode.size * 3) as usize]);
self.store(instruction.outputs[i], func(val_a, val_b, val_c, val_d));
}
}
// cargo asm "tape-drive::interpreter::IntervalInterpreter::scene" --no-color
// --rust > scene.asm
pub fn scene(
&mut self,
px: Interval,
py: Interval,
pz: Interval,
time: Interval,
) -> Interval {
self.clear_stacks();
for instruction in &self.csg.tape {
use SSAOpcode::*;
match instruction.opcode.opcode {
SSAReturn => {
return self.load(instruction.inputs[0]);
},
SSAPosition => {
self.store(instruction.outputs[0], px);
self.store(instruction.outputs[1], py);
self.store(instruction.outputs[2], pz);
self.store(instruction.outputs[3], time);
},
SSAAdd => {
self.param_two(instruction, |val_a, val_b| val_a + val_b);
},
SSASub => {
self.param_two(instruction, |val_a, val_b| val_a - val_b);
},
SSAMul => {
self.param_two(instruction, |val_a, val_b| val_a * val_b);
},
SSADiv => {
self.param_two(instruction, |val_a, val_b| val_a / val_b);
},
SSAMod => {
self.param_two(instruction, |val_a, val_b| val_a % val_b);
},
SSAAtan2 => self.param_two(instruction, |val_a, val_b| val_a.atan2(val_b)),
SSAMin => {
self.param_two(instruction, |val_a, val_b| val_a.min_choice(val_b).0);
},
SSAMinMaterial => todo!(),
SSAMax => {
self.param_two(instruction, |val_a, val_b| val_a.max_choice(val_b).0);
},
SSAMaxMaterial => todo!(),
SSADot => {
let val_a = (0..instruction.opcode.size)
.map(|i| self.load(instruction.inputs[i as usize]))
.collect::<Vec<_>>();
let val_b = (instruction.opcode.size..(instruction.opcode.size * 2))
.map(|i| self.load(instruction.inputs[i as usize]))
.collect::<Vec<_>>();
self.store(
instruction.outputs[0],
match instruction.opcode.size {
1 => val_a[0] * val_b[0],
2 => (val_a[0] * val_b[0]) + (val_a[1] * val_b[1]),
3 => {
(val_a[0] * val_b[0])
+ (val_a[1] * val_b[1])
+ (val_a[2] * val_b[2])
},
4 => {
(val_a[0] * val_b[0])
+ (val_a[1] * val_b[1])
+ (val_a[2] * val_b[2])
+ (val_a[3] * val_b[3])
},
_ => unreachable!(),
},
);
},
SSALength => {
let val_a = (0..instruction.opcode.size)
.map(|i| self.load(instruction.inputs[i as usize]))
.collect::<Vec<_>>();
let val_a_lowers = val_a
.iter()
.map(|v| {
(v.lower().simd_le(VALUE_0) & v.upper().simd_ge(VALUE_0))
.select(VALUE_0, v.lower().abs().simd_min(v.upper().abs()))
})
.collect::<Vec<_>>();
let val_a_uppers = val_a
.iter()
.map(|v| v.lower().abs().simd_max(v.upper().abs()))
.collect::<Vec<_>>();
self.store(
instruction.outputs[0],
match instruction.opcode.size {
1 => Interval::new(val_a_lowers[0], val_a_uppers[0]),
2 => Interval::new(
((val_a_lowers[0] * val_a_lowers[0])
+ (val_a_lowers[1] * val_a_lowers[1]))
.sqrt(),
((val_a_uppers[0] * val_a_uppers[0])
+ (val_a_uppers[1] * val_a_uppers[1]))
.sqrt(),
),
3 => Interval::new(
((val_a_lowers[0] * val_a_lowers[0])
+ (val_a_lowers[1] * val_a_lowers[1])
+ (val_a_lowers[2] * val_a_lowers[2]))
.sqrt(),
((val_a_uppers[0] * val_a_uppers[0])
+ (val_a_uppers[1] * val_a_uppers[1])
+ (val_a_uppers[2] * val_a_uppers[2]))
.sqrt(),
),
4 => Interval::new(
((val_a_lowers[0] * val_a_lowers[0])
+ (val_a_lowers[1] * val_a_lowers[1])
+ (val_a_lowers[2] * val_a_lowers[2])
+ (val_a_lowers[3] * val_a_lowers[3]))
.sqrt(),
((val_a_uppers[0] * val_a_uppers[0])
+ (val_a_uppers[1] * val_a_uppers[1])
+ (val_a_uppers[2] * val_a_uppers[2])
+ (val_a_uppers[3] * val_a_uppers[3]))
.sqrt(),
),
_ => unreachable!(),
},
);
},
SSADistance => {
let val_a = (0..instruction.opcode.size)
.map(|i| self.load(instruction.inputs[i as usize]))
.collect::<Vec<_>>();
let val_b = (instruction.opcode.size..(instruction.opcode.size * 2))
.map(|i| self.load(instruction.inputs[i as usize]))
.collect::<Vec<_>>();
let val_a = val_b
.into_iter()
.zip(val_a.into_iter())
.map(|(a, b)| a - b)
.collect::<Vec<_>>();
let val_a_lowers = val_a
.iter()
.map(|v| {
(v.lower().simd_le(VALUE_0) & v.upper().simd_ge(VALUE_0))
.select(VALUE_0, v.lower().abs().simd_min(v.upper().abs()))
})
.collect::<Vec<_>>();
let val_a_uppers = val_a
.iter()
.map(|v| v.lower().abs().simd_max(v.upper().abs()))
.collect::<Vec<_>>();
self.store(
instruction.outputs[0],
match instruction.opcode.size {
1 => Interval::new(val_a_lowers[0], val_a_uppers[0]),
2 => Interval::new(
((val_a_lowers[0] * val_a_lowers[0])
+ (val_a_lowers[1] * val_a_lowers[1]))
.sqrt(),
((val_a_uppers[0] * val_a_uppers[0])
+ (val_a_uppers[1] * val_a_uppers[1]))
.sqrt(),
),
3 => Interval::new(
((val_a_lowers[0] * val_a_lowers[0])
+ (val_a_lowers[1] * val_a_lowers[1])
+ (val_a_lowers[2] * val_a_lowers[2]))
.sqrt(),
((val_a_uppers[0] * val_a_uppers[0])
+ (val_a_uppers[1] * val_a_uppers[1])
+ (val_a_uppers[2] * val_a_uppers[2]))
.sqrt(),
),
4 => Interval::new(
((val_a_lowers[0] * val_a_lowers[0])
+ (val_a_lowers[1] * val_a_lowers[1])
+ (val_a_lowers[2] * val_a_lowers[2])
+ (val_a_lowers[3] * val_a_lowers[3]))
.sqrt(),
((val_a_uppers[0] * val_a_uppers[0])
+ (val_a_uppers[1] * val_a_uppers[1])
+ (val_a_uppers[2] * val_a_uppers[2])
+ (val_a_uppers[3] * val_a_uppers[3]))
.sqrt(),
),
_ => unreachable!(),
},
);
},
SSANegate => {
self.param_one(instruction, |val_a| -val_a);
},
SSARound => {
self.param_one(instruction, |val_a| val_a.round());
},
SSAAbs => {
self.param_one(instruction, |val_a| val_a.abs());
},
SSAFloor => {
self.param_one(instruction, |val_a| val_a.floor());
},
SSACeil => {
self.param_one(instruction, |val_a| val_a.ceil());
},
SSAFract => {
self.param_one(instruction, |val_a| val_a.fract());
},
SSASin => {
self.param_one(instruction, |val_a| val_a.sin());
},
SSACos => {
self.param_one(instruction, |val_a| val_a.cos());
},
SSATan => {
self.param_one(instruction, |val_a| val_a.tan());
},
SSAAsin => {
self.param_one(instruction, |val_a| val_a.asin());
},
SSAAcos => {
self.param_one(instruction, |val_a| val_a.acos());
},
SSAAtan => {
self.param_one(instruction, |val_a| val_a.atan());
},
SSAExp => {
self.param_one(instruction, |val_a| val_a.exp());
},
SSALog => {
self.param_one(instruction, |val_a| val_a.ln());
},
SSASqrt => {
self.param_one(instruction, |val_a| val_a.sqrt());
},
SSASquare => {
self.param_one(instruction, |val_a| val_a.square());
},
SSACube => {
self.param_one(instruction, |val_a| val_a.cube());
},
SSASmoothMin => {
self.param_three(instruction, |d1, d2, k| {
let h = (Interval::HALF + (Interval::HALF * (d2 - d1) / k))
.clamp(Interval::ZERO, Interval::ONE);
return ((d2 * (Interval::ONE - h)) + (d1 * h))
- k * h * (Interval::ONE - h);
});
},
SSASmoothMax => {
self.param_three(instruction, |d1, d2, k| {
let h = (Interval::HALF - (Interval::HALF * (d2 + d1) / k))
.clamp(Interval::ZERO, Interval::ONE);
return ((d2 * (Interval::ONE - h)) + (-d1 * h))
+ k * h * (Interval::ONE - h);
});
},
SSASmoothMinMaterial => todo!(),
SSASmoothMaxMaterial => todo!(),
SSAClamp => {
self.param_three(instruction, |val_a, val_b, val_c| val_a.clamp(val_b, val_c));
},
SSAMix => {
self.param_three(instruction, |val_a, val_b, val_c| {
(val_a * (Interval::ONE - val_c)) + (val_b * val_c)
});
},
SSAFMA => {
self.param_three(instruction, |val_a, val_b, val_c| (val_a * val_b) + val_c);
},
SSASDFSphere => {
let pos_x = self.load(instruction.inputs[0]);
let pos_y = self.load(instruction.inputs[1]);
let pos_z = self.load(instruction.inputs[2]);
let radius = self.load(instruction.inputs[3]);
self.store(
instruction.outputs[0],
((pos_x.square()) + (pos_y.square()) + (pos_z.square())).sqrt() - radius,
);
},
SSASDFBox => {
let pos_x = self.load(instruction.inputs[0]);
let pos_y = self.load(instruction.inputs[1]);
let pos_z = self.load(instruction.inputs[2]);
let rad_x = self.load(instruction.inputs[3]);
let rad_y = self.load(instruction.inputs[4]);
let rad_z = self.load(instruction.inputs[5]);
let qx = pos_x.abs() - rad_x;
let qy = pos_y.abs() - rad_y;
let qz = pos_z.abs() - rad_z;
let qxmax = qx.max_choice(Interval::ZERO).0;
let qymax = qy.max_choice(Interval::ZERO).0;
let qzmax = qz.max_choice(Interval::ZERO).0;
self.store(
instruction.outputs[0],
((qxmax.square()) + (qymax.square()) + (qzmax.square())).sqrt()
+ qx.max_choice(qy.max_choice(qz).0)
.0
.min_choice(Interval::ZERO)
.0,
);
},
SSASDFTorus => {
let pos_x = self.load(instruction.inputs[0]);
let pos_y = self.load(instruction.inputs[1]);
let pos_z = self.load(instruction.inputs[2]);
let radius1 = self.load(instruction.inputs[3]);
let radius2 = self.load(instruction.inputs[4]);
let q = ((pos_x.square()) + (pos_z.square())).sqrt() - radius1;
self.store(
instruction.outputs[0],
((q.square()) + (pos_y.square())).sqrt() - radius2,
);
},
SSACompare => {
self.param_two(instruction, |val_a, val_b| val_a.compare(val_b));
},
SSAAnd => {
self.param_two(instruction, |val_a, val_b| val_a.and_choice(val_b).0);
},
SSAOr => {
self.param_two(instruction, |val_a, val_b| val_a.or_choice(val_b).0);
},
SSARecip => {
self.param_one(instruction, |val_a| val_a.recip());
},
SSANot => {
self.param_one(instruction, |val_a| !val_a);
},
SSAStop => return Interval::ZERO,
}
}
return Interval::NAN;
}
}
+35
View File
@@ -0,0 +1,35 @@
use std::simd::{StdFloat, cmp::SimdPartialEq, num::SimdFloat};
use crate::BYTECODE_IMITATES_GLSL;
pub mod interval;
pub mod point;
pub type Value = std::simd::f32x8;
pub type Mask = std::simd::mask32x8;
pub const VALUE_NAN: Value = Value::splat(core::f32::NAN);
pub const VALUE_1: Value = Value::splat(1.0);
pub const VALUE_0: Value = Value::splat(0.0);
pub const VALUE_05: Value = Value::splat(0.5);
pub const VALUE_M1: Value = Value::splat(-1.0);
pub const VALUE_2: Value = Value::splat(2.0);
pub const VALUE_PI: Value = Value::splat(core::f32::consts::PI);
pub const VALUE_PI_2: Value = Value::splat(core::f32::consts::FRAC_PI_2);
pub const VALUE_TAU: Value = Value::splat(core::f32::consts::TAU);
pub fn glsign(f: Value) -> Value {
if BYTECODE_IMITATES_GLSL {
f.simd_eq(VALUE_0).select(f, f.signum())
} else {
f.signum()
}
}
pub fn glfract(f: Value) -> Value {
if BYTECODE_IMITATES_GLSL {
f - f.floor()
} else {
f.fract()
}
}
+373
View File
@@ -0,0 +1,373 @@
use std::simd::{
StdFloat,
cmp::{SimdPartialEq, SimdPartialOrd},
num::SimdFloat,
};
use crate::{
interpreters::{VALUE_0, VALUE_1, VALUE_05, VALUE_M1, VALUE_NAN, Value, glfract},
ssa::{SSAInput, SSAInstruction, SSAOpcode, SSATape},
};
#[derive(Clone, Debug)]
pub struct PointInterpreter<'csg> {
value_map: Vec<Value>,
csg: &'csg SSATape,
}
impl PointInterpreter<'_> {
fn load(&self, consider: SSAInput) -> Value {
match consider {
SSAInput::Register(r) => self.value_map[r as usize],
SSAInput::Constant(c) => Value::splat(c),
}
}
fn store(&mut self, location: u32, value: Value) {
self.value_map[location as usize] = value;
}
}
impl<'csg> PointInterpreter<'csg> {
pub fn new(csg: &'csg SSATape) -> Self {
PointInterpreter {
value_map: vec![VALUE_0; csg.last_output as usize],
csg,
}
}
fn clear_stacks(&mut self) {
self.value_map = vec![VALUE_0; self.csg.last_output as usize];
}
fn param_one(&mut self, instruction: &SSAInstruction, func: impl Fn(Value) -> Value) {
for i in 0..instruction.opcode.size as usize {
let val_a = self.load(instruction.inputs[i]);
self.store(instruction.outputs[i], func(val_a));
}
}
fn param_two(&mut self, instruction: &SSAInstruction, func: impl Fn(Value, Value) -> Value) {
for i in 0..instruction.opcode.size as usize {
let val_a = self.load(instruction.inputs[i]);
let val_b = self.load(instruction.inputs[i + instruction.opcode.size as usize]);
self.store(instruction.outputs[i], func(val_a, val_b));
}
}
fn param_three(
&mut self,
instruction: &SSAInstruction,
func: impl Fn(Value, Value, Value) -> Value,
) {
for i in 0..instruction.opcode.size as usize {
let val_a = self.load(instruction.inputs[i]);
let val_b = self.load(instruction.inputs[i + instruction.opcode.size as usize]);
let val_c = self.load(instruction.inputs[i + (instruction.opcode.size * 2) as usize]);
self.store(instruction.outputs[i], func(val_a, val_b, val_c));
}
}
fn param_four(
&mut self,
instruction: &SSAInstruction,
func: impl Fn(Value, Value, Value, Value) -> Value,
) {
for i in 0..instruction.opcode.size as usize {
let val_a = self.load(instruction.inputs[i]);
let val_b = self.load(instruction.inputs[i + instruction.opcode.size as usize]);
let val_c = self.load(instruction.inputs[i + (instruction.opcode.size * 2) as usize]);
let val_d = self.load(instruction.inputs[i + (instruction.opcode.size * 3) as usize]);
self.store(instruction.outputs[i], func(val_a, val_b, val_c, val_d));
}
}
// cargo asm "tape-drive::interpreter::PointInterpreter::scene" --no-color
// --rust > scene.asm
pub fn scene(&mut self, px: Value, py: Value, pz: Value, time: Value) -> Value {
self.clear_stacks();
for instruction in &self.csg.tape {
use SSAOpcode::*;
match instruction.opcode.opcode {
SSAReturn => {
return self.load(instruction.inputs[0]);
},
SSAPosition => {
self.store(instruction.outputs[0], px);
self.store(instruction.outputs[1], py);
self.store(instruction.outputs[2], pz);
self.store(instruction.outputs[3], time);
},
SSAAdd => {
self.param_two(instruction, |val_a, val_b| val_a + val_b);
},
SSASub => {
self.param_two(instruction, |val_a, val_b| val_a - val_b);
},
SSAMul => {
self.param_two(instruction, |val_a, val_b| val_a * val_b);
},
SSADiv => {
self.param_two(instruction, |val_a, val_b| val_a / val_b);
},
SSAMod => {
self.param_two(instruction, |val_a, val_b| val_a % val_b);
},
SSAAtan2 => self.param_two(instruction, |val_a, val_b| {
let mut val_a = val_a.to_array();
let val_b = val_b.to_array();
for i in 0..Value::LEN {
val_a[i] = val_a[i].atan2(val_b[i]);
}
Value::from_array(val_a)
}),
SSAMin => {
self.param_two(instruction, |val_a, val_b| val_a.simd_min(val_b));
},
SSAMinMaterial => todo!(),
SSAMax => {
self.param_two(instruction, |val_a, val_b| val_a.simd_max(val_b));
},
SSAMaxMaterial => todo!(),
SSADot => {
let val_a = (0..instruction.opcode.size)
.map(|i| self.load(instruction.inputs[i as usize]))
.collect::<Vec<_>>();
let val_b = (instruction.opcode.size..(instruction.opcode.size * 2))
.map(|i| self.load(instruction.inputs[i as usize]))
.collect::<Vec<_>>();
self.store(
instruction.outputs[0],
match instruction.opcode.size {
1 => val_a[0] * val_b[0],
2 => (val_a[0] * val_b[0]) + (val_a[1] * val_b[1]),
3 => {
(val_a[0] * val_b[0])
+ (val_a[1] * val_b[1])
+ (val_a[2] * val_b[2])
},
4 => {
(val_a[0] * val_b[0])
+ (val_a[1] * val_b[1])
+ (val_a[2] * val_b[2])
+ (val_a[3] * val_b[3])
},
_ => unreachable!(),
},
);
},
SSALength => {
let val_a = (0..instruction.opcode.size)
.map(|i| self.load(instruction.inputs[i as usize]))
.collect::<Vec<_>>();
self.store(
instruction.outputs[0],
match instruction.opcode.size {
1 => val_a[0],
2 => ((val_a[0] * val_a[0]) + (val_a[1] * val_a[1])).sqrt(),
3 => ((val_a[0] * val_a[0])
+ (val_a[1] * val_a[1])
+ (val_a[2] * val_a[2]))
.sqrt(),
4 => ((val_a[0] * val_a[0])
+ (val_a[1] * val_a[1])
+ (val_a[2] * val_a[2])
+ (val_a[3] * val_a[3]))
.sqrt(),
_ => unreachable!(),
},
);
},
SSADistance => {
let val_a = (0..instruction.opcode.size)
.map(|i| self.load(instruction.inputs[i as usize]))
.collect::<Vec<_>>();
let val_b = (instruction.opcode.size..(instruction.opcode.size * 2))
.map(|i| self.load(instruction.inputs[i as usize]))
.collect::<Vec<_>>();
self.store(
instruction.outputs[0],
match instruction.opcode.size {
1 => val_b[0] - val_a[0],
2 => (((val_a[0] - val_b[0]) * (val_a[0] - val_b[0]))
+ ((val_a[1] - val_b[1]) * (val_a[1] - val_b[1])))
.sqrt(),
3 => (((val_a[0] - val_b[0]) * (val_a[0] - val_b[0]))
+ ((val_a[1] - val_b[1]) * (val_a[1] - val_b[1]))
+ ((val_a[2] - val_b[2]) * (val_a[2] - val_b[2]))
+ ((val_a[3] - val_b[3]) * (val_a[3] - val_b[3])))
.sqrt(),
4 => (((val_a[0] - val_b[0]) * (val_a[0] - val_b[0]))
+ ((val_a[1] - val_b[1]) * (val_a[1] - val_b[1]))
+ ((val_a[2] - val_b[2]) * (val_a[2] - val_b[2]))
+ ((val_a[3] - val_b[3]) * (val_a[3] - val_b[3]))
+ ((val_a[4] - val_b[4]) * (val_a[4] - val_b[4])))
.sqrt(),
_ => unreachable!(),
},
);
},
SSANegate => {
self.param_one(instruction, |val_a| -val_a);
},
SSARound => {
self.param_one(instruction, |val_a| val_a.round());
},
SSAAbs => {
self.param_one(instruction, |val_a| val_a.abs());
},
SSAFloor => {
self.param_one(instruction, |val_a| val_a.floor());
},
SSACeil => {
self.param_one(instruction, |val_a| val_a.ceil());
},
SSAFract => {
self.param_one(instruction, |val_a| glfract(val_a));
},
SSASin => {
self.param_one(instruction, |val_a| val_a.sin());
},
SSACos => {
self.param_one(instruction, |val_a| val_a.cos());
},
SSATan => {
self.param_one(instruction, |val_a| {
Value::from_array(val_a.to_array().map(|f| f.tan()))
});
},
SSAAsin => {
self.param_one(instruction, |val_a| {
Value::from_array(val_a.to_array().map(|f| f.asin()))
});
},
SSAAcos => {
self.param_one(instruction, |val_a| {
Value::from_array(val_a.to_array().map(|f| f.acos()))
});
},
SSAAtan => {
self.param_one(instruction, |val_a| {
Value::from_array(val_a.to_array().map(|f| f.atan()))
});
},
SSAExp => {
self.param_one(instruction, |val_a| val_a.exp());
},
SSALog => {
self.param_one(instruction, |val_a| val_a.ln());
},
SSASqrt => {
self.param_one(instruction, |val_a| val_a.sqrt());
},
SSASquare => {
self.param_one(instruction, |val_a| val_a * val_a);
},
SSACube => {
self.param_one(instruction, |val_a| val_a * val_a * val_a);
},
SSASmoothMin => {
self.param_three(instruction, |d1, d2, k| {
let h =
(VALUE_05 + (VALUE_05 * (d2 - d1) / k)).simd_clamp(VALUE_0, VALUE_1);
return ((d2 * (VALUE_1 - h)) + (d1 * h)) - k * h * (VALUE_1 - h);
});
},
SSASmoothMax => {
self.param_three(instruction, |d1, d2, k| {
let h =
(VALUE_05 - (VALUE_05 * (d2 + d1) / k)).simd_clamp(VALUE_0, VALUE_1);
return ((d2 * (VALUE_1 - h)) + (-d1 * h)) + k * h * (VALUE_1 - h);
});
},
SSASmoothMinMaterial => todo!(),
SSASmoothMaxMaterial => todo!(),
SSAClamp => {
self.param_three(instruction, |val_a, val_b, val_c| {
val_a.simd_clamp(val_b, val_c)
});
},
SSAMix => {
self.param_three(instruction, |val_a, val_b, val_c| {
(val_a * (VALUE_1 - val_c)) + (val_b * val_c)
});
},
SSAFMA => {
self.param_three(instruction, |val_a, val_b, val_c| {
val_a.mul_add(val_b, val_c)
});
},
SSASDFSphere => {
let pos_x = self.load(instruction.inputs[0]);
let pos_y = self.load(instruction.inputs[1]);
let pos_z = self.load(instruction.inputs[2]);
let radius = self.load(instruction.inputs[3]);
self.store(
instruction.outputs[0],
((pos_x * pos_x) + (pos_y * pos_y) + (pos_z * pos_z)).sqrt() - radius,
);
},
SSASDFBox => {
let pos_x = self.load(instruction.inputs[0]);
let pos_y = self.load(instruction.inputs[1]);
let pos_z = self.load(instruction.inputs[2]);
let rad_x = self.load(instruction.inputs[3]);
let rad_y = self.load(instruction.inputs[4]);
let rad_z = self.load(instruction.inputs[5]);
let qx = pos_x.abs() - rad_x;
let qy = pos_y.abs() - rad_y;
let qz = pos_z.abs() - rad_z;
let qxmax = qx.simd_max(VALUE_0);
let qymax = qy.simd_max(VALUE_0);
let qzmax = qz.simd_max(VALUE_0);
self.store(
instruction.outputs[0],
((qxmax * qxmax) + (qymax * qymax) + (qzmax * qzmax)).sqrt()
+ qx.simd_max(qy.simd_max(qz)).simd_min(VALUE_0),
);
},
SSASDFTorus => {
let pos_x = self.load(instruction.inputs[0]);
let pos_y = self.load(instruction.inputs[1]);
let pos_z = self.load(instruction.inputs[2]);
let radius1 = self.load(instruction.inputs[3]);
let radius2 = self.load(instruction.inputs[4]);
let q = ((pos_x * pos_x) + (pos_z * pos_z)).sqrt() - radius1;
self.store(
instruction.outputs[0],
((q * q) + (pos_y * pos_y)).sqrt() - radius2,
);
},
SSACompare => {
self.param_two(instruction, |val_a, val_b| {
val_a
.simd_gt(val_b)
.select(VALUE_1, val_a.simd_lt(val_b).select(VALUE_M1, VALUE_0))
});
},
SSAAnd => {
self.param_two(instruction, |val_a, val_b| {
val_a.simd_eq(VALUE_0).select(val_a, val_b)
});
},
SSAOr => {
self.param_two(instruction, |val_a, val_b| {
val_a.simd_eq(VALUE_0).select(val_b, val_a)
});
},
SSARecip => {
self.param_one(instruction, |val_a| val_a.recip());
},
SSANot => {
self.param_one(instruction, |val_a| {
val_a.simd_eq(VALUE_0).select(VALUE_1, VALUE_0)
});
},
SSAStop => return VALUE_0,
}
}
return VALUE_NAN;
}
}
+9
View File
@@ -0,0 +1,9 @@
#![feature(portable_simd)]
pub mod instruction_set;
pub mod interpreters;
mod spirv_compilers;
pub mod ssa;
pub mod types;
pub mod vm;
const BYTECODE_IMITATES_GLSL: bool = false;
+876
View File
@@ -0,0 +1,876 @@
use foldhash::{HashMap, HashMapExt};
use rspirv::{dr::Builder, spirv};
use crate::{
spirv_compilers::SpirVTypes,
ssa::{SSAInput, SSAInstruction, SSAOpcode, SSATape},
};
pub(crate) fn compile_gradient_function(
b: &mut Builder,
tape: &SSATape,
types: SpirVTypes,
function_id: Option<spirv::Word>,
) {
let _scene = b
.begin_function(
types.float,
function_id,
//spirv::FunctionControl::DONT_INLINE
spirv::FunctionControl::INLINE
| spirv::FunctionControl::PURE
| spirv::FunctionControl::CONST,
types.point_fn_type,
)
.unwrap();
let pos_p = b.function_parameter(types.vec4p).unwrap();
b.begin_block(None).unwrap();
let pos = b.load(types.vec4, None, pos_p, None, []).unwrap();
let mut mapping = HashMap::<u32, u32>::new();
for (line, instruction) in tape.tape.iter().enumerate() {
use SSAOpcode::*;
use rspirv::dr::Operand::IdRef;
b.line(types.jit_string, line as u32, 0);
fn input_resolve(
float: u32,
b: &mut rspirv::dr::Builder,
mapping: &HashMap<u32, u32>,
value: SSAInput,
) -> u32 {
match value {
SSAInput::Register(r) => mapping[&r],
SSAInput::Constant(c) => b.constant_bit32(float, c.to_bits()),
}
}
fn param_one(
float: u32,
b: &mut rspirv::dr::Builder,
mapping: &mut HashMap<u32, u32>,
instruction: &SSAInstruction,
func: impl Fn(&mut rspirv::dr::Builder, u32) -> u32,
) {
for i in 0..instruction.opcode.size as usize {
let val_a = input_resolve(float, b, &mapping, instruction.inputs[i]);
mapping.insert(instruction.outputs[i], func(b, val_a));
}
}
fn param_two(
float: u32,
b: &mut rspirv::dr::Builder,
mapping: &mut HashMap<u32, u32>,
instruction: &SSAInstruction,
func: impl Fn(&mut rspirv::dr::Builder, u32, u32) -> u32,
) {
for i in 0..instruction.opcode.size as usize {
let val_a = input_resolve(float, b, &mapping, instruction.inputs[i]);
let val_b = input_resolve(
float,
b,
&mapping,
instruction.inputs[i + instruction.opcode.size as usize],
);
mapping.insert(instruction.outputs[i], func(b, val_a, val_b));
}
}
fn param_three(
float: u32,
b: &mut rspirv::dr::Builder,
mapping: &mut HashMap<u32, u32>,
instruction: &SSAInstruction,
func: impl Fn(&mut rspirv::dr::Builder, u32, u32, u32) -> u32,
) {
for i in 0..instruction.opcode.size as usize {
let val_a = input_resolve(float, b, &mapping, instruction.inputs[i]);
let val_b = input_resolve(
float,
b,
&mapping,
instruction.inputs[i + instruction.opcode.size as usize],
);
let val_c = input_resolve(
float,
b,
&mapping,
instruction.inputs[i + (instruction.opcode.size as usize * 2)],
);
mapping.insert(instruction.outputs[i], func(b, val_a, val_b, val_c));
}
}
fn param_four(
float: u32,
b: &mut rspirv::dr::Builder,
mapping: &mut HashMap<u32, u32>,
instruction: &SSAInstruction,
func: impl Fn(&mut rspirv::dr::Builder, u32, u32, u32, u32) -> u32,
) {
for i in 0..instruction.opcode.size as usize {
let val_a = input_resolve(float, b, &mapping, instruction.inputs[i]);
let val_b = input_resolve(
float,
b,
&mapping,
instruction.inputs[i + instruction.opcode.size as usize],
);
let val_c = input_resolve(
float,
b,
&mapping,
instruction.inputs[i + (instruction.opcode.size as usize * 2)],
);
let val_d = input_resolve(
float,
b,
&mapping,
instruction.inputs[i + (instruction.opcode.size as usize * 3)],
);
mapping.insert(instruction.outputs[i], func(b, val_a, val_b, val_c, val_d));
}
}
match instruction.opcode.opcode {
SSAStop => {
let zero = b.constant_bit32(types.float, (0.0f32).to_bits());
b.ret_value(zero).unwrap();
},
SSAReturn => {
let value = input_resolve(types.float, b, &mapping, instruction.inputs[0]);
b.ret_value(value).unwrap();
},
SSAPosition => {
mapping.insert(
instruction.outputs[0],
b.composite_extract(types.float, None, pos, [0]).unwrap(),
);
mapping.insert(
instruction.outputs[1],
b.composite_extract(types.float, None, pos, [1]).unwrap(),
);
mapping.insert(
instruction.outputs[2],
b.composite_extract(types.float, None, pos, [2]).unwrap(),
);
mapping.insert(
instruction.outputs[3],
b.composite_extract(types.float, None, pos, [3]).unwrap(),
);
},
SSAAdd => {
param_two(
types.float,
b,
&mut mapping,
instruction,
|b, val_a, val_b| b.f_add(types.float, None, val_a, val_b).unwrap(),
);
},
SSASub => {
param_two(
types.float,
b,
&mut mapping,
instruction,
|b, val_a, val_b| b.f_sub(types.float, None, val_a, val_b).unwrap(),
);
},
SSAMul => {
param_two(
types.float,
b,
&mut mapping,
instruction,
|b, val_a, val_b| b.f_mul(types.float, None, val_a, val_b).unwrap(),
);
},
SSADiv => {
param_two(
types.float,
b,
&mut mapping,
instruction,
|b, val_a, val_b| b.f_div(types.float, None, val_a, val_b).unwrap(),
);
},
SSAMod => {
param_two(
types.float,
b,
&mut mapping,
instruction,
|b, val_a, val_b| b.f_mod(types.float, None, val_a, val_b).unwrap(),
);
},
SSAAtan2 => {
param_two(
types.float,
b,
&mut mapping,
instruction,
|b, val_a, val_b| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Atan2 as u32,
[IdRef(val_a), IdRef(val_b)],
)
.unwrap()
},
);
},
SSAMin => {
param_two(
types.float,
b,
&mut mapping,
instruction,
|b, val_a, val_b| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::FMin as u32,
[IdRef(val_a), IdRef(val_b)],
)
.unwrap()
},
);
},
SSAMinMaterial => todo!(),
SSAMax => {
param_two(
types.float,
b,
&mut mapping,
instruction,
|b, val_a, val_b| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::FMax as u32,
[IdRef(val_a), IdRef(val_b)],
)
.unwrap()
},
);
},
SSAMaxMaterial => todo!(),
SSADot => {
let val_a = (0..instruction.opcode.size)
.map(|i| {
input_resolve(types.float, b, &mapping, instruction.inputs[i as usize])
})
.collect::<Vec<_>>();
let val_b = (instruction.opcode.size..(instruction.opcode.size * 2))
.map(|i| {
input_resolve(types.float, b, &mapping, instruction.inputs[i as usize])
})
.collect::<Vec<_>>();
let dot = if instruction.opcode.size == 1 {
b.f_mul(types.float, None, val_a[0], val_b[0]).unwrap()
} else {
let vector = [types.void, types.float, types.vec2, types.vec3, types.vec4]
[instruction.opcode.size as usize];
let val_a = b.composite_construct(vector, None, val_a).unwrap();
let val_b = b.composite_construct(vector, None, val_b).unwrap();
b.dot(types.float, None, val_a, val_b).unwrap()
};
mapping.insert(instruction.outputs[0], dot);
},
SSALength => {
let val_a = (0..instruction.opcode.size)
.map(|i| {
input_resolve(types.float, b, &mapping, instruction.inputs[i as usize])
})
.collect::<Vec<_>>();
let length = if instruction.opcode.size == 1 {
val_a[0]
} else {
let vector = [types.void, types.float, types.vec2, types.vec3, types.vec4]
[instruction.opcode.size as usize];
let val_a = b.composite_construct(vector, None, val_a).unwrap();
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Length as u32,
[IdRef(val_a)],
)
.unwrap()
};
mapping.insert(instruction.outputs[0], length);
},
SSADistance => {
let val_a = (0..instruction.opcode.size)
.map(|i| {
input_resolve(types.float, b, &mapping, instruction.inputs[i as usize])
})
.collect::<Vec<_>>();
let val_b = (instruction.opcode.size..(instruction.opcode.size * 2))
.map(|i| {
input_resolve(types.float, b, &mapping, instruction.inputs[i as usize])
})
.collect::<Vec<_>>();
let distance = if instruction.opcode.size == 1 {
b.f_sub(types.float, None, val_b[0], val_a[0]).unwrap()
} else {
let vector = [types.void, types.float, types.vec2, types.vec3, types.vec4]
[instruction.opcode.size as usize];
let val_a = b.composite_construct(vector, None, val_a).unwrap();
let val_b = b.composite_construct(vector, None, val_b).unwrap();
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Distance as u32,
[IdRef(val_a), IdRef(val_b)],
)
.unwrap()
};
mapping.insert(instruction.outputs[0], distance);
},
SSARecip => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
let one = b.constant_bit32(types.float, (1.0f32).to_bits());
b.f_div(types.float, None, one, val_a).unwrap()
});
},
SSANegate => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.f_negate(types.float, None, val_a).unwrap()
});
},
SSARound => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Round as u32,
[IdRef(val_a)],
)
.unwrap()
});
},
SSAAbs => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::FAbs as u32,
[IdRef(val_a)],
)
.unwrap()
});
},
SSAFloor => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Floor as u32,
[IdRef(val_a)],
)
.unwrap()
});
},
SSACeil => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Ceil as u32,
[IdRef(val_a)],
)
.unwrap()
});
},
SSAFract => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Fract as u32,
[IdRef(val_a)],
)
.unwrap()
});
},
SSASin => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Sin as u32,
[IdRef(val_a)],
)
.unwrap()
});
},
SSACos => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Cos as u32,
[IdRef(val_a)],
)
.unwrap()
});
},
SSATan => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Tan as u32,
[IdRef(val_a)],
)
.unwrap()
});
},
SSAAsin => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Asin as u32,
[IdRef(val_a)],
)
.unwrap()
});
},
SSAAcos => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Acos as u32,
[IdRef(val_a)],
)
.unwrap()
});
},
SSAAtan => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Atan as u32,
[IdRef(val_a)],
)
.unwrap()
});
},
SSAExp => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Exp as u32,
[IdRef(val_a)],
)
.unwrap()
});
},
SSALog => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Log as u32,
[IdRef(val_a)],
)
.unwrap()
});
},
SSASqrt => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Sqrt as u32,
[IdRef(val_a)],
)
.unwrap()
});
},
SSASquare => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.f_mul(types.float, None, val_a, val_a).unwrap()
});
},
SSACube => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
let square = b.f_mul(types.float, None, val_a, val_a).unwrap();
b.f_mul(types.float, None, val_a, square).unwrap()
});
},
SSASmoothMin => {
param_three(types.float, b, &mut mapping, instruction, |b, d1, d2, k| {
let half_const = b.constant_bit32(types.float, (0.5f32).to_bits());
let zero_const = b.constant_bit32(types.float, (0.0f32).to_bits());
let one_const = b.constant_bit32(types.float, (1.0f32).to_bits());
let sub = b.f_sub(types.float, None, d2, d1).unwrap();
let mul_half = b.f_mul(types.float, None, sub, half_const).unwrap();
let div_k = b.f_div(types.float, None, mul_half, k).unwrap();
let add_half = b.f_add(types.float, None, div_k, half_const).unwrap();
let h = b
.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::FClamp as u32,
[IdRef(add_half), IdRef(zero_const), IdRef(one_const)],
)
.unwrap();
let negh = b.f_sub(types.float, None, one_const, h).unwrap();
let h_negh = b.f_mul(types.float, None, h, negh).unwrap();
let kh_negh = b.f_mul(types.float, None, k, h_negh).unwrap();
let mix = b
.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::FMix as u32,
[IdRef(d2), IdRef(d1), IdRef(h)],
)
.unwrap();
b.f_sub(types.float, None, mix, kh_negh).unwrap()
});
},
SSASmoothMax => {
param_three(types.float, b, &mut mapping, instruction, |b, d1, d2, k| {
let half_const = b.constant_bit32(types.float, (0.5f32).to_bits());
let zero_const = b.constant_bit32(types.float, (0.0f32).to_bits());
let one_const = b.constant_bit32(types.float, (1.0f32).to_bits());
let sub = b.f_add(types.float, None, d2, d1).unwrap();
let mul_half = b.f_mul(types.float, None, sub, half_const).unwrap();
let div_k = b.f_div(types.float, None, mul_half, k).unwrap();
let add_half = b.f_sub(types.float, None, half_const, div_k).unwrap();
let h = b
.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::FClamp as u32,
[IdRef(add_half), IdRef(zero_const), IdRef(one_const)],
)
.unwrap();
let negh = b.f_sub(types.float, None, one_const, h).unwrap();
let h_negh = b.f_mul(types.float, None, h, negh).unwrap();
let kh_negh = b.f_mul(types.float, None, k, h_negh).unwrap();
let negate = b.f_negate(types.float, None, d1).unwrap();
let mix = b
.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::FMix as u32,
[IdRef(d2), IdRef(negate), IdRef(h)],
)
.unwrap();
b.f_add(types.float, None, mix, kh_negh).unwrap()
});
},
SSASmoothMinMaterial => todo!(),
SSASmoothMaxMaterial => todo!(),
SSAClamp => {
param_three(
types.float,
b,
&mut mapping,
instruction,
|b, val_a, val_b, val_c| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::FClamp as u32,
[IdRef(val_a), IdRef(val_b), IdRef(val_c)],
)
.unwrap()
},
);
},
SSAMix => {
param_three(
types.float,
b,
&mut mapping,
instruction,
|b, val_a, val_b, val_c| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::FMix as u32,
[IdRef(val_a), IdRef(val_b), IdRef(val_c)],
)
.unwrap()
},
);
},
SSAFMA => {
param_three(
types.float,
b,
&mut mapping,
instruction,
|b, val_a, val_b, val_c| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Fma as u32,
[IdRef(val_a), IdRef(val_b), IdRef(val_c)],
)
.unwrap()
},
);
},
SSASDFSphere => {
let pos_part = (0..3)
.map(|i| {
input_resolve(types.float, b, &mapping, instruction.inputs[i as usize])
})
.collect::<Vec<_>>();
let pos = b.composite_construct(types.vec3, None, pos_part).unwrap();
let radius =
input_resolve(types.float, b, &mapping, instruction.inputs[3 as usize]);
let length = b
.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Length as u32,
[IdRef(pos)],
)
.unwrap();
let sphere = b.f_sub(types.float, None, length, radius).unwrap();
mapping.insert(instruction.outputs[0], sphere);
},
SSASDFBox => {
let pos_part = (0..3)
.map(|i| {
input_resolve(types.float, b, &mapping, instruction.inputs[i as usize])
})
.collect::<Vec<_>>();
let pos = b.composite_construct(types.vec3, None, pos_part).unwrap();
let dim_part = (3..6)
.map(|i| {
input_resolve(types.float, b, &mapping, instruction.inputs[i as usize])
})
.collect::<Vec<_>>();
let dim = b.composite_construct(types.vec3, None, dim_part).unwrap();
let abs = b
.ext_inst(
types.vec3,
None,
types.glsl,
spirv::GLOp::FAbs as u32,
[IdRef(pos)],
)
.unwrap();
let q = b.f_sub(types.vec3, None, abs, dim).unwrap();
let zero = b.constant_bit32(types.float, 0);
let zero_vec3 = b
.composite_construct(types.vec3, None, [zero, zero, zero])
.unwrap();
let q_limit = b
.ext_inst(
types.vec3,
None,
types.glsl,
spirv::GLOp::FMax as u32,
[IdRef(q), IdRef(zero_vec3)],
)
.unwrap();
let length = b
.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Length as u32,
[IdRef(q_limit)],
)
.unwrap();
let q_x = b.composite_extract(types.float, None, q, [0]).unwrap();
let q_y = b.composite_extract(types.float, None, q, [1]).unwrap();
let q_z = b.composite_extract(types.float, None, q, [2]).unwrap();
let max1 = b
.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::FMax as u32,
[IdRef(q_x), IdRef(q_y)],
)
.unwrap();
let max2 = b
.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::FMax as u32,
[IdRef(max1), IdRef(q_z)],
)
.unwrap();
let min = b
.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::FMax as u32,
[IdRef(max2), IdRef(zero)],
)
.unwrap();
mapping.insert(
instruction.outputs[0],
b.f_add(types.float, None, length, min).unwrap(),
);
},
SSASDFTorus => {
let pos_part = (0..3)
.map(|i| {
input_resolve(types.float, b, &mapping, instruction.inputs[i as usize])
})
.collect::<Vec<_>>();
let pos_vec2 = b
.composite_construct(types.vec2, None, [pos_part[0], pos_part[2]])
.unwrap();
let rad1 = input_resolve(types.float, b, &mapping, instruction.inputs[3 as usize]);
let rad2 = input_resolve(types.float, b, &mapping, instruction.inputs[4 as usize]);
let dot = b.dot(types.float, None, pos_vec2, pos_vec2).unwrap();
let sqrt = b
.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Sqrt as u32,
[IdRef(dot)],
)
.unwrap();
let subtx = b.f_sub(types.float, None, sqrt, rad1).unwrap();
let q = b
.composite_construct(types.vec2, None, [subtx, pos_part[1]])
.unwrap();
let length = b
.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Length as u32,
[IdRef(q)],
)
.unwrap();
mapping.insert(
instruction.outputs[0],
b.f_sub(types.float, None, length, rad2).unwrap(),
);
},
SSACompare => {
param_two(
types.float,
b,
&mut mapping,
instruction,
|b, val_a, val_b| {
let nan = b.constant_bit32(types.float, f32::NAN.to_bits());
let zero = b.constant_bit32(types.float, (0.0f32).to_bits());
let one = b.constant_bit32(types.float, (1.0f32).to_bits());
let onen = b.constant_bit32(types.float, (-1.0f32).to_bits());
let equal = b.f_ord_equal(types.bool, None, val_a, val_b).unwrap();
let less = b.f_ord_less_than(types.bool, None, val_a, val_b).unwrap();
let more = b
.f_ord_greater_than(types.bool, None, val_a, val_b)
.unwrap();
let select_less = b.select(types.float, None, less, onen, nan).unwrap();
let select_more =
b.select(types.float, None, more, one, select_less).unwrap();
let select_eq = b
.select(types.float, None, equal, zero, select_more)
.unwrap();
select_eq
},
);
},
SSAAnd => {
param_two(
types.float,
b,
&mut mapping,
instruction,
|b, val_a, val_b| {
let zero = b.constant_bit32(types.float, (0.0f32).to_bits());
let equal = b.f_ord_equal(types.bool, None, val_a, zero).unwrap();
b.select(types.float, None, equal, val_a, val_b).unwrap()
},
);
},
SSAOr => {
param_two(
types.float,
b,
&mut mapping,
instruction,
|b, val_a, val_b| {
let zero = b.constant_bit32(types.float, (0.0f32).to_bits());
let equal = b.f_ord_equal(types.bool, None, val_a, zero).unwrap();
b.select(types.float, None, equal, val_b, val_a).unwrap()
},
);
},
SSANot => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
let zero = b.constant_bit32(types.float, (0.0f32).to_bits());
let one = b.constant_bit32(types.float, (1.0f32).to_bits());
let equal = b.f_ord_equal(types.bool, None, val_a, zero).unwrap();
b.select(types.float, None, equal, one, zero).unwrap()
});
},
}
}
b.end_function().unwrap();
}
File diff suppressed because it is too large Load Diff
+18
View File
@@ -0,0 +1,18 @@
pub(crate) mod gradient;
pub(crate) mod point;
//pub(crate) mod interval;
#[derive(Debug, Clone, Copy)]
pub(crate) struct SpirVTypes {
pub glsl: u32,
pub void: u32,
pub float: u32,
pub bool: u32,
pub vec2: u32,
pub vec3: u32,
pub vec4: u32,
pub vec4p: u32,
pub point_fn_type: u32,
pub interval_fn_type: u32,
pub jit_string: u32,
}
+876
View File
@@ -0,0 +1,876 @@
use foldhash::{HashMap, HashMapExt};
use rspirv::{dr::Builder, spirv};
use crate::{
spirv_compilers::SpirVTypes,
ssa::{SSAInput, SSAInstruction, SSAOpcode, SSATape},
};
pub(crate) fn compile_point_function(
b: &mut Builder,
tape: &SSATape,
types: SpirVTypes,
function_id: Option<spirv::Word>,
) {
let _scene = b
.begin_function(
types.float,
function_id,
//spirv::FunctionControl::DONT_INLINE
spirv::FunctionControl::INLINE
| spirv::FunctionControl::PURE
| spirv::FunctionControl::CONST,
types.point_fn_type,
)
.unwrap();
let pos_p = b.function_parameter(types.vec4p).unwrap();
b.begin_block(None).unwrap();
let pos = b.load(types.vec4, None, pos_p, None, []).unwrap();
let mut mapping = HashMap::<u32, u32>::new();
for (line, instruction) in tape.tape.iter().enumerate() {
use SSAOpcode::*;
use rspirv::dr::Operand::IdRef;
b.line(types.jit_string, line as u32, 0);
fn input_resolve(
float: u32,
b: &mut rspirv::dr::Builder,
mapping: &HashMap<u32, u32>,
value: SSAInput,
) -> u32 {
match value {
SSAInput::Register(r) => mapping[&r],
SSAInput::Constant(c) => b.constant_bit32(float, c.to_bits()),
}
}
fn param_one(
float: u32,
b: &mut rspirv::dr::Builder,
mapping: &mut HashMap<u32, u32>,
instruction: &SSAInstruction,
func: impl Fn(&mut rspirv::dr::Builder, u32) -> u32,
) {
for i in 0..instruction.opcode.size as usize {
let val_a = input_resolve(float, b, &mapping, instruction.inputs[i]);
mapping.insert(instruction.outputs[i], func(b, val_a));
}
}
fn param_two(
float: u32,
b: &mut rspirv::dr::Builder,
mapping: &mut HashMap<u32, u32>,
instruction: &SSAInstruction,
func: impl Fn(&mut rspirv::dr::Builder, u32, u32) -> u32,
) {
for i in 0..instruction.opcode.size as usize {
let val_a = input_resolve(float, b, &mapping, instruction.inputs[i]);
let val_b = input_resolve(
float,
b,
&mapping,
instruction.inputs[i + instruction.opcode.size as usize],
);
mapping.insert(instruction.outputs[i], func(b, val_a, val_b));
}
}
fn param_three(
float: u32,
b: &mut rspirv::dr::Builder,
mapping: &mut HashMap<u32, u32>,
instruction: &SSAInstruction,
func: impl Fn(&mut rspirv::dr::Builder, u32, u32, u32) -> u32,
) {
for i in 0..instruction.opcode.size as usize {
let val_a = input_resolve(float, b, &mapping, instruction.inputs[i]);
let val_b = input_resolve(
float,
b,
&mapping,
instruction.inputs[i + instruction.opcode.size as usize],
);
let val_c = input_resolve(
float,
b,
&mapping,
instruction.inputs[i + (instruction.opcode.size as usize * 2)],
);
mapping.insert(instruction.outputs[i], func(b, val_a, val_b, val_c));
}
}
fn param_four(
float: u32,
b: &mut rspirv::dr::Builder,
mapping: &mut HashMap<u32, u32>,
instruction: &SSAInstruction,
func: impl Fn(&mut rspirv::dr::Builder, u32, u32, u32, u32) -> u32,
) {
for i in 0..instruction.opcode.size as usize {
let val_a = input_resolve(float, b, &mapping, instruction.inputs[i]);
let val_b = input_resolve(
float,
b,
&mapping,
instruction.inputs[i + instruction.opcode.size as usize],
);
let val_c = input_resolve(
float,
b,
&mapping,
instruction.inputs[i + (instruction.opcode.size as usize * 2)],
);
let val_d = input_resolve(
float,
b,
&mapping,
instruction.inputs[i + (instruction.opcode.size as usize * 3)],
);
mapping.insert(instruction.outputs[i], func(b, val_a, val_b, val_c, val_d));
}
}
match instruction.opcode.opcode {
SSAStop => {
let zero = b.constant_bit32(types.float, (0.0f32).to_bits());
b.ret_value(zero).unwrap();
},
SSAReturn => {
let value = input_resolve(types.float, b, &mapping, instruction.inputs[0]);
b.ret_value(value).unwrap();
},
SSAPosition => {
mapping.insert(
instruction.outputs[0],
b.composite_extract(types.float, None, pos, [0]).unwrap(),
);
mapping.insert(
instruction.outputs[1],
b.composite_extract(types.float, None, pos, [1]).unwrap(),
);
mapping.insert(
instruction.outputs[2],
b.composite_extract(types.float, None, pos, [2]).unwrap(),
);
mapping.insert(
instruction.outputs[3],
b.composite_extract(types.float, None, pos, [3]).unwrap(),
);
},
SSAAdd => {
param_two(
types.float,
b,
&mut mapping,
instruction,
|b, val_a, val_b| b.f_add(types.float, None, val_a, val_b).unwrap(),
);
},
SSASub => {
param_two(
types.float,
b,
&mut mapping,
instruction,
|b, val_a, val_b| b.f_sub(types.float, None, val_a, val_b).unwrap(),
);
},
SSAMul => {
param_two(
types.float,
b,
&mut mapping,
instruction,
|b, val_a, val_b| b.f_mul(types.float, None, val_a, val_b).unwrap(),
);
},
SSADiv => {
param_two(
types.float,
b,
&mut mapping,
instruction,
|b, val_a, val_b| b.f_div(types.float, None, val_a, val_b).unwrap(),
);
},
SSAMod => {
param_two(
types.float,
b,
&mut mapping,
instruction,
|b, val_a, val_b| b.f_mod(types.float, None, val_a, val_b).unwrap(),
);
},
SSAAtan2 => {
param_two(
types.float,
b,
&mut mapping,
instruction,
|b, val_a, val_b| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Atan2 as u32,
[IdRef(val_a), IdRef(val_b)],
)
.unwrap()
},
);
},
SSAMin => {
param_two(
types.float,
b,
&mut mapping,
instruction,
|b, val_a, val_b| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::FMin as u32,
[IdRef(val_a), IdRef(val_b)],
)
.unwrap()
},
);
},
SSAMinMaterial => todo!(),
SSAMax => {
param_two(
types.float,
b,
&mut mapping,
instruction,
|b, val_a, val_b| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::FMax as u32,
[IdRef(val_a), IdRef(val_b)],
)
.unwrap()
},
);
},
SSAMaxMaterial => todo!(),
SSADot => {
let val_a = (0..instruction.opcode.size)
.map(|i| {
input_resolve(types.float, b, &mapping, instruction.inputs[i as usize])
})
.collect::<Vec<_>>();
let val_b = (instruction.opcode.size..(instruction.opcode.size * 2))
.map(|i| {
input_resolve(types.float, b, &mapping, instruction.inputs[i as usize])
})
.collect::<Vec<_>>();
let dot = if instruction.opcode.size == 1 {
b.f_mul(types.float, None, val_a[0], val_b[0]).unwrap()
} else {
let vector = [types.void, types.float, types.vec2, types.vec3, types.vec4]
[instruction.opcode.size as usize];
let val_a = b.composite_construct(vector, None, val_a).unwrap();
let val_b = b.composite_construct(vector, None, val_b).unwrap();
b.dot(types.float, None, val_a, val_b).unwrap()
};
mapping.insert(instruction.outputs[0], dot);
},
SSALength => {
let val_a = (0..instruction.opcode.size)
.map(|i| {
input_resolve(types.float, b, &mapping, instruction.inputs[i as usize])
})
.collect::<Vec<_>>();
let length = if instruction.opcode.size == 1 {
val_a[0]
} else {
let vector = [types.void, types.float, types.vec2, types.vec3, types.vec4]
[instruction.opcode.size as usize];
let val_a = b.composite_construct(vector, None, val_a).unwrap();
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Length as u32,
[IdRef(val_a)],
)
.unwrap()
};
mapping.insert(instruction.outputs[0], length);
},
SSADistance => {
let val_a = (0..instruction.opcode.size)
.map(|i| {
input_resolve(types.float, b, &mapping, instruction.inputs[i as usize])
})
.collect::<Vec<_>>();
let val_b = (instruction.opcode.size..(instruction.opcode.size * 2))
.map(|i| {
input_resolve(types.float, b, &mapping, instruction.inputs[i as usize])
})
.collect::<Vec<_>>();
let distance = if instruction.opcode.size == 1 {
b.f_sub(types.float, None, val_b[0], val_a[0]).unwrap()
} else {
let vector = [types.void, types.float, types.vec2, types.vec3, types.vec4]
[instruction.opcode.size as usize];
let val_a = b.composite_construct(vector, None, val_a).unwrap();
let val_b = b.composite_construct(vector, None, val_b).unwrap();
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Distance as u32,
[IdRef(val_a), IdRef(val_b)],
)
.unwrap()
};
mapping.insert(instruction.outputs[0], distance);
},
SSARecip => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
let one = b.constant_bit32(types.float, (1.0f32).to_bits());
b.f_div(types.float, None, one, val_a).unwrap()
});
},
SSANegate => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.f_negate(types.float, None, val_a).unwrap()
});
},
SSARound => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Round as u32,
[IdRef(val_a)],
)
.unwrap()
});
},
SSAAbs => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::FAbs as u32,
[IdRef(val_a)],
)
.unwrap()
});
},
SSAFloor => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Floor as u32,
[IdRef(val_a)],
)
.unwrap()
});
},
SSACeil => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Ceil as u32,
[IdRef(val_a)],
)
.unwrap()
});
},
SSAFract => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Fract as u32,
[IdRef(val_a)],
)
.unwrap()
});
},
SSASin => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Sin as u32,
[IdRef(val_a)],
)
.unwrap()
});
},
SSACos => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Cos as u32,
[IdRef(val_a)],
)
.unwrap()
});
},
SSATan => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Tan as u32,
[IdRef(val_a)],
)
.unwrap()
});
},
SSAAsin => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Asin as u32,
[IdRef(val_a)],
)
.unwrap()
});
},
SSAAcos => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Acos as u32,
[IdRef(val_a)],
)
.unwrap()
});
},
SSAAtan => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Atan as u32,
[IdRef(val_a)],
)
.unwrap()
});
},
SSAExp => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Exp as u32,
[IdRef(val_a)],
)
.unwrap()
});
},
SSALog => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Log as u32,
[IdRef(val_a)],
)
.unwrap()
});
},
SSASqrt => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Sqrt as u32,
[IdRef(val_a)],
)
.unwrap()
});
},
SSASquare => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
b.f_mul(types.float, None, val_a, val_a).unwrap()
});
},
SSACube => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
let square = b.f_mul(types.float, None, val_a, val_a).unwrap();
b.f_mul(types.float, None, val_a, square).unwrap()
});
},
SSASmoothMin => {
param_three(types.float, b, &mut mapping, instruction, |b, d1, d2, k| {
let half_const = b.constant_bit32(types.float, (0.5f32).to_bits());
let zero_const = b.constant_bit32(types.float, (0.0f32).to_bits());
let one_const = b.constant_bit32(types.float, (1.0f32).to_bits());
let sub = b.f_sub(types.float, None, d2, d1).unwrap();
let mul_half = b.f_mul(types.float, None, sub, half_const).unwrap();
let div_k = b.f_div(types.float, None, mul_half, k).unwrap();
let add_half = b.f_add(types.float, None, div_k, half_const).unwrap();
let h = b
.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::FClamp as u32,
[IdRef(add_half), IdRef(zero_const), IdRef(one_const)],
)
.unwrap();
let negh = b.f_sub(types.float, None, one_const, h).unwrap();
let h_negh = b.f_mul(types.float, None, h, negh).unwrap();
let kh_negh = b.f_mul(types.float, None, k, h_negh).unwrap();
let mix = b
.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::FMix as u32,
[IdRef(d2), IdRef(d1), IdRef(h)],
)
.unwrap();
b.f_sub(types.float, None, mix, kh_negh).unwrap()
});
},
SSASmoothMax => {
param_three(types.float, b, &mut mapping, instruction, |b, d1, d2, k| {
let half_const = b.constant_bit32(types.float, (0.5f32).to_bits());
let zero_const = b.constant_bit32(types.float, (0.0f32).to_bits());
let one_const = b.constant_bit32(types.float, (1.0f32).to_bits());
let sub = b.f_add(types.float, None, d2, d1).unwrap();
let mul_half = b.f_mul(types.float, None, sub, half_const).unwrap();
let div_k = b.f_div(types.float, None, mul_half, k).unwrap();
let add_half = b.f_sub(types.float, None, half_const, div_k).unwrap();
let h = b
.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::FClamp as u32,
[IdRef(add_half), IdRef(zero_const), IdRef(one_const)],
)
.unwrap();
let negh = b.f_sub(types.float, None, one_const, h).unwrap();
let h_negh = b.f_mul(types.float, None, h, negh).unwrap();
let kh_negh = b.f_mul(types.float, None, k, h_negh).unwrap();
let negate = b.f_negate(types.float, None, d1).unwrap();
let mix = b
.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::FMix as u32,
[IdRef(d2), IdRef(negate), IdRef(h)],
)
.unwrap();
b.f_add(types.float, None, mix, kh_negh).unwrap()
});
},
SSASmoothMinMaterial => todo!(),
SSASmoothMaxMaterial => todo!(),
SSAClamp => {
param_three(
types.float,
b,
&mut mapping,
instruction,
|b, val_a, val_b, val_c| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::FClamp as u32,
[IdRef(val_a), IdRef(val_b), IdRef(val_c)],
)
.unwrap()
},
);
},
SSAMix => {
param_three(
types.float,
b,
&mut mapping,
instruction,
|b, val_a, val_b, val_c| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::FMix as u32,
[IdRef(val_a), IdRef(val_b), IdRef(val_c)],
)
.unwrap()
},
);
},
SSAFMA => {
param_three(
types.float,
b,
&mut mapping,
instruction,
|b, val_a, val_b, val_c| {
b.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Fma as u32,
[IdRef(val_a), IdRef(val_b), IdRef(val_c)],
)
.unwrap()
},
);
},
SSASDFSphere => {
let pos_part = (0..3)
.map(|i| {
input_resolve(types.float, b, &mapping, instruction.inputs[i as usize])
})
.collect::<Vec<_>>();
let pos = b.composite_construct(types.vec3, None, pos_part).unwrap();
let radius =
input_resolve(types.float, b, &mapping, instruction.inputs[3 as usize]);
let length = b
.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Length as u32,
[IdRef(pos)],
)
.unwrap();
let sphere = b.f_sub(types.float, None, length, radius).unwrap();
mapping.insert(instruction.outputs[0], sphere);
},
SSASDFBox => {
let pos_part = (0..3)
.map(|i| {
input_resolve(types.float, b, &mapping, instruction.inputs[i as usize])
})
.collect::<Vec<_>>();
let pos = b.composite_construct(types.vec3, None, pos_part).unwrap();
let dim_part = (3..6)
.map(|i| {
input_resolve(types.float, b, &mapping, instruction.inputs[i as usize])
})
.collect::<Vec<_>>();
let dim = b.composite_construct(types.vec3, None, dim_part).unwrap();
let abs = b
.ext_inst(
types.vec3,
None,
types.glsl,
spirv::GLOp::FAbs as u32,
[IdRef(pos)],
)
.unwrap();
let q = b.f_sub(types.vec3, None, abs, dim).unwrap();
let zero = b.constant_bit32(types.float, 0);
let zero_vec3 = b
.composite_construct(types.vec3, None, [zero, zero, zero])
.unwrap();
let q_limit = b
.ext_inst(
types.vec3,
None,
types.glsl,
spirv::GLOp::FMax as u32,
[IdRef(q), IdRef(zero_vec3)],
)
.unwrap();
let length = b
.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Length as u32,
[IdRef(q_limit)],
)
.unwrap();
let q_x = b.composite_extract(types.float, None, q, [0]).unwrap();
let q_y = b.composite_extract(types.float, None, q, [1]).unwrap();
let q_z = b.composite_extract(types.float, None, q, [2]).unwrap();
let max1 = b
.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::FMax as u32,
[IdRef(q_x), IdRef(q_y)],
)
.unwrap();
let max2 = b
.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::FMax as u32,
[IdRef(max1), IdRef(q_z)],
)
.unwrap();
let min = b
.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::FMax as u32,
[IdRef(max2), IdRef(zero)],
)
.unwrap();
mapping.insert(
instruction.outputs[0],
b.f_add(types.float, None, length, min).unwrap(),
);
},
SSASDFTorus => {
let pos_part = (0..3)
.map(|i| {
input_resolve(types.float, b, &mapping, instruction.inputs[i as usize])
})
.collect::<Vec<_>>();
let pos_vec2 = b
.composite_construct(types.vec2, None, [pos_part[0], pos_part[2]])
.unwrap();
let rad1 = input_resolve(types.float, b, &mapping, instruction.inputs[3 as usize]);
let rad2 = input_resolve(types.float, b, &mapping, instruction.inputs[4 as usize]);
let dot = b.dot(types.float, None, pos_vec2, pos_vec2).unwrap();
let sqrt = b
.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Sqrt as u32,
[IdRef(dot)],
)
.unwrap();
let subtx = b.f_sub(types.float, None, sqrt, rad1).unwrap();
let q = b
.composite_construct(types.vec2, None, [subtx, pos_part[1]])
.unwrap();
let length = b
.ext_inst(
types.float,
None,
types.glsl,
spirv::GLOp::Length as u32,
[IdRef(q)],
)
.unwrap();
mapping.insert(
instruction.outputs[0],
b.f_sub(types.float, None, length, rad2).unwrap(),
);
},
SSACompare => {
param_two(
types.float,
b,
&mut mapping,
instruction,
|b, val_a, val_b| {
let nan = b.constant_bit32(types.float, f32::NAN.to_bits());
let zero = b.constant_bit32(types.float, (0.0f32).to_bits());
let one = b.constant_bit32(types.float, (1.0f32).to_bits());
let onen = b.constant_bit32(types.float, (-1.0f32).to_bits());
let equal = b.f_ord_equal(types.bool, None, val_a, val_b).unwrap();
let less = b.f_ord_less_than(types.bool, None, val_a, val_b).unwrap();
let more = b
.f_ord_greater_than(types.bool, None, val_a, val_b)
.unwrap();
let select_less = b.select(types.float, None, less, onen, nan).unwrap();
let select_more =
b.select(types.float, None, more, one, select_less).unwrap();
let select_eq = b
.select(types.float, None, equal, zero, select_more)
.unwrap();
select_eq
},
);
},
SSAAnd => {
param_two(
types.float,
b,
&mut mapping,
instruction,
|b, val_a, val_b| {
let zero = b.constant_bit32(types.float, (0.0f32).to_bits());
let equal = b.f_ord_equal(types.bool, None, val_a, zero).unwrap();
b.select(types.float, None, equal, val_a, val_b).unwrap()
},
);
},
SSAOr => {
param_two(
types.float,
b,
&mut mapping,
instruction,
|b, val_a, val_b| {
let zero = b.constant_bit32(types.float, (0.0f32).to_bits());
let equal = b.f_ord_equal(types.bool, None, val_a, zero).unwrap();
b.select(types.float, None, equal, val_b, val_a).unwrap()
},
);
},
SSANot => {
param_one(types.float, b, &mut mapping, instruction, |b, val_a| {
let zero = b.constant_bit32(types.float, (0.0f32).to_bits());
let one = b.constant_bit32(types.float, (1.0f32).to_bits());
let equal = b.f_ord_equal(types.bool, None, val_a, zero).unwrap();
b.select(types.float, None, equal, one, zero).unwrap()
});
},
}
}
b.end_function().unwrap();
}
+536
View File
@@ -0,0 +1,536 @@
use core::f32;
use rspirv::{binary::Disassemble, dr::Module, spirv};
use crate::{
instruction_set::InstructionSet,
spirv_compilers::{
SpirVTypes,
//gradient::compile_gradient_function,
//interval::compile_interval_function,
point::compile_point_function,
},
};
const JIT_VERSION: u32 = 1;
#[derive(Debug, Default, PartialEq, Eq, Clone, Copy)]
pub enum SSAOpcode {
SSAAdd,
SSASub,
SSAMul,
SSADiv,
SSAAtan2,
SSAMin,
SSAMax,
SSACompare,
SSAMod,
SSAAnd,
SSAOr,
SSANegate,
SSAAbs,
SSARecip,
SSASqrt,
SSASquare,
SSAFloor,
SSACeil,
SSARound,
SSASin,
SSACos,
SSATan,
SSAAsin,
SSAAcos,
SSAAtan,
SSAExp,
SSALog,
SSANot,
SSAFract,
SSACube,
SSASmoothMin,
SSASmoothMax,
SSAClamp,
SSAMix,
SSAFMA,
SSADot,
SSALength,
SSADistance,
#[default]
SSAStop,
SSAReturn,
SSAPosition,
SSAMinMaterial,
SSAMaxMaterial,
SSASmoothMinMaterial,
SSASmoothMaxMaterial,
SSASDFSphere,
SSASDFBox,
SSASDFTorus,
}
#[derive(Debug, Default, PartialEq, Eq, Clone, Copy)]
pub struct SSAOpcodeSized {
pub opcode: SSAOpcode,
pub size: u8,
}
#[derive(Debug, PartialEq, Clone, Copy)]
pub enum SSAInput {
Register(u32),
Constant(f32),
}
#[derive(Debug, Default, PartialEq, Clone)]
pub struct SSAInstruction {
pub opcode: SSAOpcodeSized,
pub inputs: Vec<SSAInput>,
pub outputs: Vec<u32>,
}
#[derive(Debug, Default, PartialEq, Clone)]
pub struct SSATape {
pub last_output: u32,
pub tape: Vec<SSAInstruction>,
constants: Vec<f32>,
}
#[derive(Debug, Default, PartialEq, Eq, Clone, Copy)]
struct GPUOpcode(u8);
pub struct GPUTape {
pub instructions: Vec<u8>,
pub io: Vec<u8>,
pub constants: Vec<f32>,
}
impl SSAOpcodeSized {
const fn output(&self) -> u8 {
use SSAOpcode::*;
match self.opcode {
SSAStop => 0,
SSAReturn => 0,
SSAPosition => 4,
SSAMinMaterial => 1,
SSAMaxMaterial => 1,
SSASmoothMinMaterial => 1,
SSASmoothMaxMaterial => 1,
SSADistance => 1,
SSALength => 1,
SSADot => 1,
SSASDFSphere => 1,
SSASDFBox => 1,
SSASDFTorus => 1,
_ => self.size,
}
}
const fn input(&self) -> u8 {
use SSAOpcode::*;
match self.opcode {
SSAStop => 0,
SSAReturn => 1,
SSAPosition => 0,
SSAMinMaterial => 2,
SSAMaxMaterial => 2,
SSASmoothMinMaterial => 3,
SSASmoothMaxMaterial => 3,
SSASDFSphere => 3 + 1,
SSASDFBox => 3 + 3,
SSASDFTorus => 3 + 2,
SSAAdd | SSASub | SSAMul | SSADiv | SSAAtan2 | SSAMin | SSAMax | SSACompare
| SSAMod | SSAAnd | SSAOr | SSADot | SSADistance => self.size * 2,
SSASmoothMin | SSASmoothMax | SSAClamp | SSAMix | SSAFMA => self.size * 3,
_ => self.size,
}
}
const fn lifetime_elementwise(&self) -> (u8, u8) {
use SSAOpcode::*;
match self.opcode {
SSAStop | SSAReturn | SSAPosition | SSAMinMaterial | SSAMaxMaterial
| SSASmoothMinMaterial | SSASmoothMaxMaterial | SSASDFSphere | SSASDFBox
| SSASDFTorus | SSADot | SSADistance | SSALength => (0, 0),
SSAAdd | SSASub | SSAMul | SSADiv | SSAAtan2 | SSAMin | SSACompare | SSAMod
| SSAAnd | SSAOr => (2, 1),
SSASmoothMin | SSASmoothMax | SSAClamp | SSAMix | SSAFMA => (3, 1),
_ => (1, 1),
}
}
const fn to_raw_opcode(&self) -> GPUOpcode {
use InstructionSet::*;
use SSAOpcode::*;
const fn opcode_drop(inst: InstructionSet, width: u8) -> GPUOpcode {
assert!(width <= 4);
assert!(width > 0);
GPUOpcode(inst as u8 + ((width - 1) << 6))
}
match self.opcode {
SSAStop => opcode_drop(OPReturn, 1),
SSAReturn => opcode_drop(OPReturn, 1),
SSAPosition => opcode_drop(OPPosition, 1),
SSAMinMaterial => opcode_drop(OPMinMaterial, 1),
SSAMaxMaterial => opcode_drop(OPMaxMaterial, 1),
SSASmoothMinMaterial => opcode_drop(OPSmoothMinMaterial, 1),
SSASmoothMaxMaterial => opcode_drop(OPSmoothMaxMaterial, 1),
SSASDFSphere => opcode_drop(OPSDFSphere, 1),
SSASDFBox => opcode_drop(OPSDFBox, 1),
SSASDFTorus => opcode_drop(OPSDFTorus, 1),
SSAAdd => opcode_drop(OPAdd, self.size),
SSASub => opcode_drop(OPSub, self.size),
SSAMul => opcode_drop(OPMul, self.size),
SSADiv => opcode_drop(OPDiv, self.size),
SSAMod => opcode_drop(OPMod, self.size),
SSAAtan2 => opcode_drop(OPAtan2, self.size),
SSAMin => opcode_drop(OPMin, self.size),
SSAMax => opcode_drop(OPMax, self.size),
SSADot => opcode_drop(OPDot, self.size),
SSALength => opcode_drop(OPLength, self.size),
SSADistance => opcode_drop(OPDistance, self.size),
SSANegate => opcode_drop(OPNegate, self.size),
SSARound => opcode_drop(OPRound, self.size),
SSAAbs => opcode_drop(OPAbs, self.size),
SSAFloor => opcode_drop(OPFloor, self.size),
SSACeil => opcode_drop(OPCeil, self.size),
SSAFract => opcode_drop(OPFract, self.size),
SSASin => opcode_drop(OPSin, self.size),
SSACos => opcode_drop(OPCos, self.size),
SSATan => opcode_drop(OPTan, self.size),
SSAAsin => opcode_drop(OPAsin, self.size),
SSAAcos => opcode_drop(OPAcos, self.size),
SSAAtan => opcode_drop(OPAtan, self.size),
SSAExp => opcode_drop(OPExp, self.size),
SSALog => opcode_drop(OPLog, self.size),
SSASqrt => opcode_drop(OPSqrt, self.size),
SSASquare => opcode_drop(OPSquare, self.size),
SSACube => opcode_drop(OPCube, self.size),
SSASmoothMin => opcode_drop(OPSmoothMin, self.size),
SSASmoothMax => opcode_drop(OPSmoothMax, self.size),
SSAClamp => opcode_drop(OPClamp, self.size),
SSAMix => opcode_drop(OPMix, self.size),
SSAFMA => opcode_drop(OPFMA, self.size),
SSACompare => opcode_drop(OPCompare, self.size),
SSAAnd => opcode_drop(OPAnd, self.size),
SSAOr => opcode_drop(OPOr, self.size),
SSARecip => opcode_drop(OPRecip, self.size),
SSANot => opcode_drop(OPNot, self.size),
}
}
}
impl SSATape {
pub fn push_instruction(
&mut self,
opcode: SSAOpcodeSized,
inputs: Vec<SSAInput>,
) -> Vec<SSAInput> {
assert!(
inputs
.iter()
.filter_map(|i| match i {
SSAInput::Constant(_) => None,
SSAInput::Register(r) => Some(r),
})
.all(|&i| i < self.last_output)
);
assert!(opcode.size <= 4);
assert!(opcode.size >= 1);
assert_eq!(inputs.len(), opcode.input() as usize);
let outputs = (self.last_output..)
.take(opcode.output() as usize)
.collect::<Vec<u32>>();
let outputs_register = outputs
.iter()
.map(|r| SSAInput::Register(*r))
.collect::<Vec<SSAInput>>();
self.last_output += opcode.output() as u32;
self.constants.extend(inputs.iter().filter_map(|i| match i {
SSAInput::Constant(0.0) => None,
SSAInput::Constant(c) => Some(c),
SSAInput::Register(_) => None,
}));
self.tape.push(SSAInstruction {
opcode,
inputs,
outputs,
});
return outputs_register;
}
pub fn compile_to_gpu(&self) -> GPUTape {
let mut lifetimes = Vec::<(u32, u32)>::with_capacity(self.last_output as usize);
let mut time_unit = 0;
for SSAInstruction {
opcode,
inputs,
outputs,
} in self.tape.iter()
{
let per_element = opcode.lifetime_elementwise();
if per_element == (0, 0) {
for &value in inputs {
if let SSAInput::Register(r) = value {
assert!((r as usize) < lifetimes.len());
lifetimes[r as usize].1 = time_unit;
}
}
for &value in outputs {
assert_eq!(value as usize, lifetimes.len());
lifetimes.push((time_unit, time_unit));
}
time_unit += 1;
} else {
let mut input_iterators = (0..per_element.0)
.map(|i| inputs.iter().skip(i.into()).step_by(per_element.0.into()))
.collect::<Vec<_>>();
let mut output_iterators = (0..per_element.1)
.map(|i| outputs.iter().skip(i.into()).step_by(per_element.1.into()))
.collect::<Vec<_>>();
assert_eq!(
opcode.input() / per_element.0,
opcode.output() / per_element.1
);
for _ in 0..(opcode.input() / per_element.0) {
for iterator in input_iterators.iter_mut() {
let &value = iterator.next().unwrap();
if let SSAInput::Register(r) = value {
assert!((r as usize) < lifetimes.len());
lifetimes[r as usize].1 = time_unit;
}
}
for iterator in output_iterators.iter_mut() {
let &value = iterator.next().unwrap();
assert_eq!(value as usize, lifetimes.len());
lifetimes.push((time_unit, time_unit));
}
time_unit += 1;
}
}
}
// Registers are held UNTIL, non inclusive.
let mut register_hold = [0u32; 14];
let mut register_allocation = vec![0u8; lifetimes.len()];
for ((life_start, life_end), allocation) in
lifetimes.iter().zip(register_allocation.iter_mut())
{
if life_start == life_end {
*allocation = 0;
} else if let Some(register) = register_hold.iter().position(|reg| reg <= life_start) {
register_hold[register] = *life_end;
*allocation = (register + 1) as u8;
} else {
panic!("Failed to allocate registers");
}
}
let mut gpu_tape = GPUTape {
instructions: vec![],
io: vec![],
constants: self.constants.clone(),
};
let mut low_nibble = true;
let mut staging_byte = 0u8;
for SSAInstruction {
opcode,
inputs,
outputs,
} in self.tape.iter()
{
let code = opcode.to_raw_opcode();
gpu_tape.instructions.push(code.0);
let per_element = opcode.lifetime_elementwise();
if per_element == (0, 0) {
for input in inputs {
let register = match input {
SSAInput::Constant(0.0) => 0,
SSAInput::Constant(_) => 15,
SSAInput::Register(u) => register_allocation[*u as usize],
};
if low_nibble {
staging_byte |= register;
} else {
staging_byte |= register << 4;
gpu_tape.io.push(staging_byte);
}
low_nibble = !low_nibble;
}
for output in outputs {
let register = register_allocation[*output as usize];
if low_nibble {
staging_byte |= register;
} else {
staging_byte |= register << 4;
gpu_tape.io.push(staging_byte);
}
low_nibble = !low_nibble;
}
// Stop is implemented as SSAReturn(0);
if opcode.opcode == SSAOpcode::SSAStop {
let register = 0;
if low_nibble {
staging_byte |= register;
} else {
staging_byte |= register << 4;
gpu_tape.io.push(staging_byte);
}
low_nibble = !low_nibble;
}
} else {
let mut input_iterators = (0..per_element.0)
.map(|i| inputs.iter().skip(i.into()).step_by(per_element.0.into()))
.collect::<Vec<_>>();
let mut output_iterators = (0..per_element.1)
.map(|i| outputs.iter().skip(i.into()).step_by(per_element.1.into()))
.collect::<Vec<_>>();
assert_eq!(
opcode.input() / per_element.0,
opcode.output() / per_element.1
);
for _ in 0..(opcode.input() / per_element.0) {
for iterator in input_iterators.iter_mut() {
let &value = iterator.next().unwrap();
let register = match value {
SSAInput::Constant(0.0) => 0,
SSAInput::Constant(_) => 15,
SSAInput::Register(u) => register_allocation[u as usize],
};
if low_nibble {
staging_byte |= register;
} else {
staging_byte |= register << 4;
gpu_tape.io.push(staging_byte);
}
low_nibble = !low_nibble;
}
for iterator in output_iterators.iter_mut() {
let &value = iterator.next().unwrap();
let register = register_allocation[value as usize];
if low_nibble {
staging_byte |= register;
} else {
staging_byte |= register << 4;
gpu_tape.io.push(staging_byte);
}
low_nibble = !low_nibble;
}
}
}
}
if !low_nibble {
gpu_tape.io.push(staging_byte);
}
gpu_tape
}
pub fn compile_to_spirv(&self, module: Option<Module>) -> rspirv::dr::Module {
let with_module = module.is_some();
let mut b = if let Some(module) = module {
rspirv::dr::Builder::new_from_module(module)
} else {
let mut b = rspirv::dr::Builder::new();
b.set_version(1, 6);
b.module_processed(format!("Tape Drive JIT {JIT_VERSION}"));
b.memory_model(spirv::AddressingModel::Logical, spirv::MemoryModel::GLSL450);
b
};
let glsl = if with_module {
1
} else {
b.ext_inst_import("GLSL.std.450")
};
let void = b.type_void();
let float = b.type_float(32);
let bool = b.type_bool();
let vec2 = b.type_vector(float, 2);
let vec3 = b.type_vector(float, 3);
let vec4 = b.type_vector(float, 4);
let vec4p = b.type_pointer(None, spirv::StorageClass::Function, vec4);
let point_fn_type = b.type_function(float, vec![vec4p]);
let interval_fn_type = b.type_function(vec2, vec![vec4p, vec4p]);
let jit_string = b.string("JIT");
let types = SpirVTypes {
glsl,
void,
float,
bool,
vec2,
vec3,
vec4,
vec4p,
point_fn_type,
interval_fn_type,
jit_string,
};
compile_point_function(
&mut b,
&self,
types,
if with_module { Some(1000) } else { None },
);
//compile_interval_function(
// &mut b,
// &self,
// types,
// if with_module { Some(2000) } else { None },
//);
//compile_gradient_function(
// &mut b,
// &self,
// types,
// if with_module { Some(3000) } else { None },
//);
let module = b.module();
std::fs::write(
format!(
"{}.spv-dis",
humantime::format_rfc3339(std::time::SystemTime::now())
)
.to_string(),
module.disassemble(),
)
.unwrap();
module
}
}
+365
View File
@@ -0,0 +1,365 @@
use glam::Vec4;
/// A point in space with associated partial derivatives.
#[derive(Copy, Clone, Debug, Default, PartialEq)]
#[repr(C)]
pub struct Grad {
/// Value of the distance field at this point
pub v: f32,
/// Partial derivative with respect to `x`
pub dx: f32,
/// Partial derivative with respect to `y`
pub dy: f32,
/// Partial derivative with respect to `z`
pub dz: f32,
}
impl std::fmt::Display for Grad {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "({}, {}, {}, {})", self.v, self.dx, self.dy, self.dz)
}
}
impl Grad {
/// Constructs a new gradient
#[inline]
pub fn new(v: f32, dx: f32, dy: f32, dz: f32) -> Self {
Self { v, dx, dy, dz }
}
/// Looks up a gradient by index (0 = x, 1 = y, 2 = z)
///
/// # Panics
/// If the index is not in the 0-2 range
#[inline]
pub fn d(&self, i: usize) -> f32 {
match i {
0 => self.dx,
1 => self.dy,
2 => self.dz,
_ => panic!("invalid index {i}"),
}
}
/// Absolute value
#[inline]
pub fn abs(self) -> Self {
if self.v < 0.0 {
Grad {
v: -self.v,
dx: -self.dx,
dy: -self.dy,
dz: -self.dz,
}
} else {
self
}
}
/// Square root
#[inline]
pub fn sqrt(self) -> Self {
let v = self.v.sqrt();
Grad {
v,
dx: self.dx / (2.0 * v),
dy: self.dy / (2.0 * v),
dz: self.dz / (2.0 * v),
}
}
/// Sine
#[inline]
pub fn sin(self) -> Self {
let c = self.v.cos();
Grad {
v: self.v.sin(),
dx: self.dx * c,
dy: self.dy * c,
dz: self.dz * c,
}
}
/// Cosine
#[inline]
pub fn cos(self) -> Self {
let s = -self.v.sin();
Grad {
v: self.v.cos(),
dx: self.dx * s,
dy: self.dy * s,
dz: self.dz * s,
}
}
/// Tangent
#[inline]
pub fn tan(self) -> Self {
let c = self.v.cos().powi(2);
Grad {
v: self.v.tan(),
dx: self.dx / c,
dy: self.dy / c,
dz: self.dz / c,
}
}
/// Arcsin
#[inline]
pub fn asin(self) -> Self {
let r = (1.0 - self.v.powi(2)).sqrt();
Grad {
v: self.v.asin(),
dx: self.dx / r,
dy: self.dy / r,
dz: self.dz / r,
}
}
/// Arccos
#[inline]
pub fn acos(self) -> Self {
let r = (1.0 - self.v.powi(2)).sqrt();
Grad {
v: self.v.acos(),
dx: -self.dx / r,
dy: -self.dy / r,
dz: -self.dz / r,
}
}
/// Arctangent
#[inline]
pub fn atan(self) -> Self {
let r = self.v.powi(2) + 1.0;
Grad {
v: self.v.atan(),
dx: self.dx / r,
dy: self.dy / r,
dz: self.dz / r,
}
}
/// Exponential function
#[inline]
pub fn exp(self) -> Self {
let v = self.v.exp();
Grad {
v,
dx: v * self.dx,
dy: v * self.dy,
dz: v * self.dz,
}
}
/// Natural log
#[inline]
pub fn ln(self) -> Self {
Grad {
v: self.v.ln(),
dx: self.dx / self.v,
dy: self.dy / self.v,
dz: self.dz / self.v,
}
}
/// Reciprocal
#[inline]
pub fn recip(self) -> Self {
let v2 = -self.v.powi(2);
Grad {
v: 1.0 / self.v,
dx: self.dx / v2,
dy: self.dy / v2,
dz: self.dz / v2,
}
}
/// Minimum of two values
#[inline]
pub fn min(self, rhs: Self) -> Self {
if self.v < rhs.v { self } else { rhs }
}
/// Maximum of two values
#[inline]
pub fn max(self, rhs: Self) -> Self {
if self.v > rhs.v { self } else { rhs }
}
/// Least non-negative remainder
#[inline]
pub fn rem_euclid(&self, rhs: Grad) -> Self {
let e = self.v.div_euclid(rhs.v);
Grad {
v: self.v.rem_euclid(rhs.v),
dx: self.dx - rhs.dx * e,
dy: self.dy - rhs.dy * e,
dz: self.dz - rhs.dz * e,
}
}
/// Snap to the largest less-than-or-equal value
#[inline]
pub fn floor(&self) -> Self {
Grad {
v: self.v.floor(),
dx: 0.0,
dy: 0.0,
dz: 0.0,
}
}
/// Snap to the smallest greater-than-or-equal value
#[inline]
pub fn ceil(&self) -> Self {
Grad {
v: self.v.ceil(),
dx: 0.0,
dy: 0.0,
dz: 0.0,
}
}
/// Rounds to the nearest integer
#[inline]
pub fn round(&self) -> Self {
Grad {
v: self.v.round(),
dx: 0.0,
dy: 0.0,
dz: 0.0,
}
}
/// Four-quadrant arctangent
#[inline]
pub fn atan2(self, x: Self) -> Self {
let y = self;
let d = x.v.powi(2) + y.v.powi(2);
Grad {
v: y.v.atan2(x.v),
dx: (x.v * y.dx - y.v * x.dx) / d,
dy: (x.v * y.dy - y.v * x.dy) / d,
dz: (x.v * y.dz - y.v * x.dz) / d,
}
}
/// Checks that the two values are roughly equal, panicking otherwise
#[cfg(test)]
pub(crate) fn compare_eq(&self, other: Self) {
let d = (self.v - other.v)
.abs()
.max((self.dx - other.dx).abs())
.max((self.dy - other.dy).abs())
.max((self.dz - other.dz).abs());
if d >= 1e-6 {
panic!("lhs != rhs ({self:?} != {other:?})");
}
}
}
impl From<f32> for Grad {
#[inline]
fn from(v: f32) -> Self {
Grad {
v,
dx: 0.0,
dy: 0.0,
dz: 0.0,
}
}
}
impl From<Grad> for Vec4 {
#[inline]
fn from(g: Grad) -> Self {
Vec4::new(g.dx, g.dy, g.dz, g.v)
}
}
impl std::ops::Add<Grad> for Grad {
type Output = Self;
#[inline]
fn add(self, rhs: Self) -> Self {
Grad {
v: self.v + rhs.v,
dx: self.dx + rhs.dx,
dy: self.dy + rhs.dy,
dz: self.dz + rhs.dz,
}
}
}
impl std::ops::Mul<Grad> for Grad {
type Output = Self;
#[inline]
fn mul(self, rhs: Self) -> Self {
Self {
v: self.v * rhs.v,
dx: self.v * rhs.dx + rhs.v * self.dx,
dy: self.v * rhs.dy + rhs.v * self.dy,
dz: self.v * rhs.dz + rhs.v * self.dz,
}
}
}
impl std::ops::Mul<f32> for Grad {
type Output = Self;
#[inline]
fn mul(self, rhs: f32) -> Self {
Self {
v: self.v * rhs,
dx: self.dx * rhs,
dy: self.dy * rhs,
dz: self.dz * rhs,
}
}
}
impl std::ops::Div<Grad> for Grad {
type Output = Self;
#[inline]
fn div(self, rhs: Self) -> Self {
let d = rhs.v.powi(2);
Self {
v: self.v / rhs.v,
dx: (rhs.v * self.dx - self.v * rhs.dx) / d,
dy: (rhs.v * self.dy - self.v * rhs.dy) / d,
dz: (rhs.v * self.dz - self.v * rhs.dz) / d,
}
}
}
impl std::ops::Sub<Grad> for Grad {
type Output = Self;
#[inline]
fn sub(self, rhs: Self) -> Self {
Self {
v: self.v - rhs.v,
dx: self.dx - rhs.dx,
dy: self.dy - rhs.dy,
dz: self.dz - rhs.dz,
}
}
}
impl std::ops::Neg for Grad {
type Output = Self;
#[inline]
fn neg(self) -> Self {
Self {
v: -self.v,
dx: -self.dx,
dy: -self.dy,
dz: -self.dz,
}
}
}
+712
View File
@@ -0,0 +1,712 @@
use std::simd::{
StdFloat,
cmp::{SimdPartialEq, SimdPartialOrd},
num::SimdFloat,
};
use crate::{
interpreters::{
Mask, VALUE_0, VALUE_1, VALUE_2, VALUE_05, VALUE_M1, VALUE_NAN, VALUE_PI, VALUE_PI_2,
VALUE_TAU, Value, glfract, glsign,
},
vm::choice::{Choice, VChoice},
};
/// Stores a range, with conservative calculations to guarantee that it always
/// contains the actual value.
///
/// # Warning
/// This implementation does not set rounding modes, so it may not be _perfect_.
#[derive(Copy, Clone, PartialEq)]
#[repr(C)]
pub struct Interval {
lower: Value,
upper: Value,
}
impl std::fmt::Debug for Interval {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> Result<(), std::fmt::Error> {
f.debug_tuple("")
.field(&self.lower)
.field(&self.upper)
.finish()
}
}
impl Interval {
pub const HALF: Self = Self::const_splat(0.5);
pub const MONE: Self = Self::const_splat(-1.0);
pub const NAN: Self = Self::const_splat(core::f32::NAN);
pub const ONE: Self = Self::const_splat(1.0);
pub const PI: Self = Self::const_splat(core::f32::consts::PI);
pub const PI_2: Self = Self::const_splat(core::f32::consts::FRAC_PI_2);
pub const ZERO: Self = Self::const_splat(0.0);
/// Builds a new interval
///
/// There are two kinds of valid interval:
/// - `[lower, upper]` where `lower <= upper`
/// - `[NaN, NaN]`
///
/// # Panics
/// Panics if the resulting interval would be invalid
#[inline]
pub fn new(lower: Value, upper: Value) -> Self {
assert!(
(upper.simd_ge(lower) | (lower.is_nan() & upper.is_nan())).all(),
"invalid interval [{lower:?}, {upper:?}]"
);
Self { lower, upper }
}
#[inline]
pub const fn new_unchecked(lower: Value, upper: Value) -> Self {
Self { lower, upper }
}
pub fn splat(value: f32) -> Interval {
Interval::new(Value::splat(value), Value::splat(value))
}
pub fn splat2(lower: f32, upper: f32) -> Interval {
Interval::new(Value::splat(lower), Value::splat(upper))
}
pub const fn const_splat(value: f32) -> Interval {
Interval {
lower: Value::splat(value),
upper: Value::splat(value),
}
}
pub const fn const_splat2(lower: f32, upper: f32) -> Interval {
Interval {
lower: Value::splat(lower),
upper: Value::splat(upper),
}
}
/// Returns the lower bound of the interval
#[inline]
pub fn lower(&self) -> Value {
self.lower
}
/// Returns the upper bound of the interval
#[inline]
pub fn upper(&self) -> Value {
self.upper
}
/// Checks whether the given value is (strictly) contained in the interval
#[inline]
pub fn contains(&self, v: Value) -> Mask {
v.simd_ge(self.lower) & v.simd_le(self.upper)
}
/// Returns `true` if either bound of the interval is `NaN`
#[inline]
pub fn has_nan(&self) -> Mask {
self.lower.is_nan() | self.upper.is_nan()
}
/// Calculates the absolute value of the interval
#[inline]
pub fn abs(self) -> Self {
let llt = self.lower.simd_lt(VALUE_0);
let ugt = self.upper.simd_gt(VALUE_0);
let lower = llt.select(ugt.select(VALUE_0, -self.upper), self.lower);
let upper = llt.select(
ugt.select(self.upper.simd_max(-self.lower), -self.lower),
self.upper,
);
Interval::new(lower, upper)
}
/// Squares 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 square(self) -> Self {
let ult = self.upper.simd_lt(VALUE_0);
let lgt = self.lower.simd_gt(VALUE_0);
let has_nan = self.has_nan();
let lower = ult.select(
self.upper * self.upper,
lgt.select(self.lower * self.lower, has_nan.select(VALUE_NAN, VALUE_0)),
);
let upper = ult.select(
self.lower * self.lower,
lgt.select(
self.upper * self.upper,
has_nan.select(VALUE_NAN, {
let k = self.lower.abs().simd_max(self.upper.abs());
k * k
}),
),
);
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 {
(self - Self::PI_2).cos()
}
/// Computes the cosine of the interval
#[inline]
pub fn cos(self) -> Self {
let lower_cycle = (self.lower / VALUE_PI).floor();
let upper_cycle = (self.upper / VALUE_PI).floor();
let same_cycle = lower_cycle.simd_eq(upper_cycle);
let cycle = upper_cycle % VALUE_2;
let within_one_cycle = (upper_cycle - lower_cycle).simd_eq(VALUE_1);
let temp0 = self.lower.cos();
let temp1 = self.upper.cos();
let lower = self.has_nan().select(
VALUE_NAN,
(same_cycle | (within_one_cycle & (cycle.simd_eq(VALUE_0))))
.select(temp0.simd_min(temp1), VALUE_M1),
);
let upper = self.has_nan().select(
VALUE_NAN,
(same_cycle | (within_one_cycle & (cycle.simd_eq(VALUE_1))))
.select(temp0.simd_max(temp1), VALUE_1),
);
Interval::new(lower, upper)
}
/// Computes the tangent of the interval
///
/// Returns the `NAN` interval if the result contains a undefined point
#[inline]
pub fn tan(self) -> Self {
let size = self.upper - self.lower;
let lower_tmp = Value::from_array(self.lower.to_array().map(|f| f.tan()));
let upper_tmp = Value::from_array(self.upper.to_array().map(|f| f.tan()));
let lower =
(size.simd_lt(VALUE_PI) & upper_tmp.simd_ge(lower_tmp)).select(lower_tmp, VALUE_NAN);
let upper =
(size.simd_lt(VALUE_PI) & upper_tmp.simd_ge(lower_tmp)).select(upper_tmp, VALUE_NAN);
Interval::new(lower, upper)
}
/// Computes the arcsine of the interval
///
/// Returns the `NAN` interval if the input is invalid
#[inline]
pub fn asin(self) -> Self {
let lower = (self.lower.simd_lt(VALUE_M1) | self.upper.simd_gt(VALUE_1)).select(
VALUE_NAN,
Value::from_array(self.lower.to_array().map(|f| f.asin())),
);
let upper = (self.lower.simd_lt(VALUE_M1) | self.upper.simd_gt(VALUE_1)).select(
VALUE_NAN,
Value::from_array(self.upper.to_array().map(|f| f.asin())),
);
Interval::new(lower, upper)
}
/// Computes the arccosine of the interval
///
/// Returns the `NAN` interval if the input is invalid
#[inline]
pub fn acos(self) -> Self {
let lower = (self.lower.simd_lt(VALUE_M1) | self.upper.simd_gt(VALUE_1)).select(
VALUE_NAN,
Value::from_array(self.upper.to_array().map(|f| f.asin())),
);
let upper = (self.lower.simd_lt(VALUE_M1) | self.upper.simd_gt(VALUE_1)).select(
VALUE_NAN,
Value::from_array(self.lower.to_array().map(|f| f.asin())),
);
Interval::new(lower, upper)
}
/// Computes the arctangent of the interval
#[inline]
pub fn atan(self) -> Self {
let lower = Value::from_array(self.lower.to_array().map(|f| f.asin()));
let upper = Value::from_array(self.upper.to_array().map(|f| f.asin()));
Interval::new(lower, upper)
}
/// Computes the exponent function applied to the interval
#[inline]
pub fn exp(self) -> Self {
Interval::new(self.lower.exp(), self.upper.exp())
}
/// Computes the natural log of the input interval
///
/// Returns the `NAN` interval if the input contains zero
#[inline]
pub fn ln(self) -> Self {
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
///
/// If the interval contains values below 0, returns a `NAN` interval.
#[inline]
pub fn sqrt(self) -> Self {
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
///
/// If the interval includes 0, returns the `NAN` interval
#[inline]
pub fn recip(self) -> Self {
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 [`VChoice`] indicating whether one side is
/// always less than the other.
///
/// If either side is `NAN`, returns the `NAN` interval and `VChoice::Both`.
#[inline]
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(
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 [`VChoice`] indicating whether one side is
/// always greater than the other.
///
/// If either side is `NAN`, returns the `NAN` interval and `VChoice::Both`.
#[inline]
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(
has_nan.select(VALUE_NAN, self.lower.simd_max(rhs.lower)),
has_nan.select(VALUE_NAN, self.upper.simd_max(rhs.upper)),
),
VChoice(choice),
)
}
/// Calculates the short-circuiting `AND` of two intervals
///
/// 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, 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 [`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, 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) -> Value {
(self.lower + self.upper) / VALUE_2
}
/// Splits the interval at the midpoint
///
/// ```
/// # use fidget::types::Interval;
/// let a = Interval::new(0.0, 1.0);
/// let (lo, hi) = a.split();
/// assert_eq!(lo, Interval::new(0.0, 0.5));
/// assert_eq!(hi, Interval::new(0.5, 1.0));
/// ```
#[inline]
pub fn split(self) -> (Self, Self) {
let mid = self.midpoint();
(
Interval::new(self.lower, mid),
Interval::new(mid, self.upper),
)
}
/// Linear interpolation from `lower` to `upper`
///
/// ```
/// # use fidget::types::Interval;
/// let a = Interval::new(0.0, 2.0);
/// assert_eq!(a.lerp(0.5), 1.0);
/// assert_eq!(a.lerp(0.75), 1.5);
/// assert_eq!(a.lerp(2.0), 4.0);
/// ```
#[inline]
pub fn lerp(self, frac: Value) -> Value {
self.lower * (VALUE_1 - frac) + self.upper * frac
}
/// Calculates the width of the interval
///
/// ```
/// # use fidget::types::Interval;
/// let a = Interval::new(2.0, 3.0);
/// assert_eq!(a.width(), 1.0);
/// let b = Interval::new(2.0, 5.0);
/// assert_eq!(b.width(), 3.0);
/// ```
#[inline]
pub fn width(self) -> Value {
self.upper - self.lower
}
/// Checks that the two values are roughly equal, panicking otherwise
#[cfg(test)]
pub(crate) fn compare_eq(&self, other: Self) {
let d = (self.lower - other.lower)
.abs()
.simd_max((self.upper - other.upper).abs());
if d.simd_ge(Value::splat(1e-6)).any() {
panic!("lhs != rhs ({self:?} != {other:?})");
}
}
/// Largest value that is less-than-or-equal to this value
#[inline]
pub fn floor(&self) -> Self {
Interval::new(self.lower.floor(), self.upper.floor())
}
/// Smallest value that is greater-than-or-equal to this value
#[inline]
pub fn ceil(&self) -> Self {
Interval::new(self.lower.ceil(), self.upper.ceil())
}
/// Rounded value
#[inline]
pub fn round(&self) -> Self {
Interval::new(self.lower.round(), self.upper.round())
}
/// Four-quadrant arctangent
#[inline]
pub fn atan2(self, x: Self) -> Self {
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),
)
}
#[inline]
pub fn sign(self) -> Interval {
Interval::new(glsign(self.lower), glsign(self.upper))
}
#[inline]
pub fn fract(self) -> Interval {
let ge1 = (self.upper - self.lower).simd_ge(VALUE_1);
Interval::new(
ge1.select(VALUE_0, glfract(self.lower)),
ge1.select(VALUE_1, glfract(self.upper)),
)
}
pub fn compare(self, other: Interval) -> Interval {
let has_nan = self.has_nan() | other.has_nan();
let check1 = self.upper.simd_lt(other.lower);
let check2 = other.upper.simd_lt(self.lower);
let check3 = other.upper.simd_eq(self.lower);
let check4 = self.upper.simd_eq(other.lower);
let lower = has_nan.select(
VALUE_NAN,
check1.select(
VALUE_M1,
check2.select(
VALUE_1,
check3.select(VALUE_0, check4.select(VALUE_M1, VALUE_M1)),
),
),
);
let upper = has_nan.select(
VALUE_NAN,
check1.select(
VALUE_M1,
check2.select(
VALUE_1,
check3.select(VALUE_1, check4.select(VALUE_0, VALUE_1)),
),
),
);
Interval::new(lower, upper)
}
pub fn clamp(self, min: Interval, max: Interval) -> Interval {
Interval::new(
self.lower.simd_clamp(min.lower, max.lower),
self.upper.simd_clamp(min.upper, max.upper),
)
}
}
impl std::fmt::Display for Interval {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "({:?}, {:?})", self.lower, self.upper)
}
}
impl From<[Value; 2]> for Interval {
#[inline]
fn from(i: [Value; 2]) -> Interval {
Interval::new(i[0], i[1])
}
}
impl From<Value> for Interval {
#[inline]
fn from(f: Value) -> Self {
Interval::new(f, f)
}
}
impl std::ops::Not for Interval {
type Output = Self;
fn not(self) -> Self::Output {
let has_nan = self.has_nan();
let is_zero = self.lower.simd_eq(VALUE_0) & self.upper.simd_eq(VALUE_0);
let crosses_zero = self.lower.simd_le(VALUE_0) & self.upper.simd_ge(VALUE_0);
let lower = has_nan.select(VALUE_NAN, is_zero.select(VALUE_1, VALUE_0));
let upper = has_nan.select(VALUE_NAN, crosses_zero.select(VALUE_1, VALUE_0));
Interval::new(lower, upper)
}
}
impl std::ops::Rem<Interval> for Interval {
type Output = Self;
#[inline]
fn rem(self, rhs: Interval) -> Self::Output {
// TODO optimize this more?
let has_nan = self.has_nan() | rhs.has_nan() | rhs.contains(VALUE_0);
let other_constant = rhs.lower.simd_eq(rhs.upper) & rhs.lower.simd_gt(VALUE_0);
let a = self.lower / rhs.lower;
let b = self.upper / rhs.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 % rhs.lower, VALUE_0),
);
let upper = has_nan.select(
VALUE_NAN,
(other_constant & floors).select(self.upper % rhs.lower, rhs.abs().upper),
);
Interval::new(lower, upper)
}
}
impl std::ops::Add<Interval> for Interval {
type Output = Self;
#[inline]
fn add(self, rhs: Self) -> Self {
Interval::new(self.lower + rhs.lower, self.upper + rhs.upper)
}
}
impl std::ops::Mul<Interval> for Interval {
type Output = Self;
#[inline]
fn mul(self, rhs: Self) -> Self {
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] {
out[k] = i * j;
k += 1;
}
}
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),
)
}
}
impl std::ops::Mul<Value> for Interval {
type Output = Self;
#[inline]
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)),
)
}
}
impl std::ops::Div<Interval> for Interval {
type Output = Self;
#[inline]
fn div(self, rhs: Self) -> Self {
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.simd_min(v);
upper = upper.simd_max(v);
}
Interval::new(
has_nan.select(VALUE_NAN, lower),
has_nan.select(VALUE_NAN, upper),
)
}
}
impl std::ops::Sub<Interval> for Interval {
type Output = Self;
#[inline]
fn sub(self, rhs: Self) -> Self {
Interval::new(self.lower - rhs.upper, self.upper - rhs.lower)
}
}
impl std::ops::Neg for Interval {
type Output = Self;
#[inline]
fn neg(self) -> Self {
Interval::new(-self.upper, -self.lower)
}
}
+6
View File
@@ -0,0 +1,6 @@
//! Custom types used during evaluation
mod grad;
mod interval;
pub use grad::Grad;
pub use interval::Interval;
+96
View File
@@ -0,0 +1,96 @@
use std::simd::{
StdFloat,
cmp::{SimdPartialEq, SimdPartialOrd},
num::SimdFloat,
u32x8,
};
/// 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!(),
}
}
}
pub struct VChoice(pub 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))
}
}
+418
View File
@@ -0,0 +1,418 @@
//! General-purpose tapes for use during evaluation or further compilation
use crate::{
Error,
compiler::{RegOp, RegTape, RegisterAllocator, SsaOp, SsaTape},
context::{Context, Node},
var::VarMap,
vm::Choice,
};
use serde::{Deserialize, Serialize};
use std::sync::Arc;
/// A flattened math expression, ready for evaluation or further compilation.
///
/// Under the hood, [`VmData`] stores two different representations:
/// - A tape in [single static assignment form](https://en.wikipedia.org/wiki/Static_single-assignment_form)
/// ([`SsaTape`]), which is suitable for use during tape simplification
/// - A tape in register-allocated form ([`RegTape`]), which can be efficiently
/// evaluated or lowered into machine assembly
///
/// # Example
/// Consider the expression `x + y`. The SSA tape will look something like
/// this:
/// ```text
/// $0 = INPUT 0 // X
/// $1 = INPUT 1 // Y
/// $2 = ADD $0 $1 // (X + Y)
/// ```
///
/// This will be lowered into a tape using real (or VM) registers:
/// ```text
/// r0 = INPUT 0 // X
/// r1 = INPUT 1 // Y
/// r0 = ADD r0 r1 // (X + Y)
/// ```
///
/// Note that in this form, registers are reused (e.g. `r0` stores both `X` and
/// `X + Y`).
///
/// We can peek at the internals and see this register-allocated tape:
/// ```
/// use fidget::{
/// compiler::RegOp,
/// context::{Context, Tree},
/// vm::VmData,
/// var::Var,
/// };
///
/// let tree = Tree::x() + Tree::y();
/// let mut ctx = Context::new();
/// let sum = ctx.import(&tree);
/// let data = VmData::<255>::new(&ctx, &[sum])?;
/// assert_eq!(data.len(), 4); // X, Y, (X + Y), and output
///
/// let mut iter = data.iter_asm();
/// let vars = &data.vars; // map from var to index
/// assert_eq!(iter.next().unwrap(), RegOp::Input(0, vars[&Var::X] as u32));
/// assert_eq!(iter.next().unwrap(), RegOp::Input(1, vars[&Var::Y] as u32));
/// assert_eq!(iter.next().unwrap(), RegOp::AddRegReg(0, 0, 1));
/// # Ok::<(), fidget::Error>(())
/// ```
///
/// Despite this peek at its internals, users are unlikely to touch `VmData`
/// directly; a [`VmShape`](crate::vm::VmShape) wraps the `VmData` and
/// implements our common traits.
#[derive(Default, Serialize, Deserialize)]
pub struct VmData<const N: usize = { u8::MAX as usize }> {
ssa: SsaTape,
asm: RegTape,
/// Mapping from variables to indices during evaluation
///
/// This member is stored in a shared pointer because it's passed down to
/// children (constructed with [`VmData::simplify`]).
pub vars: Arc<VarMap>,
}
impl<const N: usize> VmData<N> {
/// Builds a new tape for the given node
pub fn new(context: &Context, nodes: &[Node]) -> Result<Self, Error> {
let (ssa, vars) = SsaTape::new(context, nodes)?;
let asm = RegTape::new::<N>(&ssa);
Ok(Self {
ssa,
asm,
vars: vars.into(),
})
}
/// Returns the length of the internal VM tape
pub fn len(&self) -> usize {
self.asm.len()
}
/// Returns true if the internal VM tape is empty
pub fn is_empty(&self) -> bool {
self.asm.is_empty()
}
/// Returns the number of choice (min/max) nodes in the tape.
///
/// This is required because some evaluators pre-allocate spaces for the
/// choice array.
pub fn choice_count(&self) -> usize {
self.ssa.choice_count
}
/// Returns the number of output nodes in the tape.
///
/// This is required because some evaluators pre-allocate spaces for the
/// output array.
pub fn output_count(&self) -> usize {
self.ssa.output_count
}
/// Returns the number of slots used by the inner VM tape
pub fn slot_count(&self) -> usize {
self.asm.slot_count()
}
/// Simplifies both inner tapes, using the provided choice array
///
/// To minimize allocations, this function takes a [`VmWorkspace`] and
/// spare [`VmData`]; it will reuse those allocations.
pub fn simplify<const M: usize>(
&self,
choices: &[Choice],
workspace: &mut VmWorkspace<M>,
mut tape: VmData<M>,
) -> Result<VmData<M>, Error> {
if choices.len() != self.choice_count() {
return Err(Error::BadChoiceSlice(
choices.len(),
self.choice_count(),
));
}
tape.ssa.reset();
// Steal `tape.asm` and hand it to the workspace for use in allocator
workspace.reset(self.ssa.tape.len(), tape.asm);
let mut choice_count = 0;
let mut output_count = 0;
// Other iterators to consume various arrays in order
let mut choice_iter = choices.iter().rev();
let mut ops_out = tape.ssa.tape;
for mut op in self.ssa.tape.iter().cloned() {
let index = match &mut op {
SsaOp::Output(reg, _i) => {
*reg = workspace.get_or_insert_active(*reg);
workspace.alloc.op(op);
ops_out.push(op);
output_count += 1;
continue;
}
_ => op.output().unwrap(),
};
if workspace.active(index).is_none() {
if op.has_choice() {
choice_iter.next().unwrap();
}
continue;
}
// Because we reassign nodes when they're used as an *input*
// (while walking the tape in reverse), this node must have been
// assigned already.
let new_index = workspace.active(index).unwrap();
match &mut op {
SsaOp::Output(..) => unreachable!(),
SsaOp::Input(index, ..) | SsaOp::CopyImm(index, ..) => {
*index = new_index;
}
SsaOp::NegReg(index, arg)
| SsaOp::AbsReg(index, arg)
| SsaOp::RecipReg(index, arg)
| SsaOp::SqrtReg(index, arg)
| SsaOp::SquareReg(index, arg)
| SsaOp::FloorReg(index, arg)
| SsaOp::CeilReg(index, arg)
| SsaOp::RoundReg(index, arg)
| SsaOp::SinReg(index, arg)
| SsaOp::CosReg(index, arg)
| SsaOp::TanReg(index, arg)
| SsaOp::AsinReg(index, arg)
| SsaOp::AcosReg(index, arg)
| SsaOp::AtanReg(index, arg)
| SsaOp::ExpReg(index, arg)
| SsaOp::LnReg(index, arg)
| SsaOp::NotReg(index, arg) => {
*index = new_index;
*arg = workspace.get_or_insert_active(*arg);
}
SsaOp::CopyReg(index, src) => {
// CopyReg effectively does
// dst <= src
// If src has not yet been used (as we iterate backwards
// through the tape), then we can replace it with dst
// everywhere!
match workspace.active(*src) {
Some(new_src) => {
*index = new_index;
*src = new_src;
}
None => {
workspace.set_active(*src, new_index);
continue;
}
}
}
SsaOp::MinRegImm(index, arg, imm)
| SsaOp::MaxRegImm(index, arg, imm)
| SsaOp::AndRegImm(index, arg, imm)
| SsaOp::OrRegImm(index, arg, imm) => {
match choice_iter.next().unwrap() {
Choice::Left => match workspace.active(*arg) {
Some(new_arg) => {
op = SsaOp::CopyReg(new_index, new_arg);
}
None => {
workspace.set_active(*arg, new_index);
continue;
}
},
Choice::Right => {
op = SsaOp::CopyImm(new_index, *imm);
}
Choice::Both => {
choice_count += 1;
*index = new_index;
*arg = workspace.get_or_insert_active(*arg);
}
Choice::Unknown => panic!("oh no"),
}
}
SsaOp::MinRegReg(index, lhs, rhs)
| SsaOp::MaxRegReg(index, lhs, rhs)
| SsaOp::AndRegReg(index, lhs, rhs)
| SsaOp::OrRegReg(index, lhs, rhs) => {
match choice_iter.next().unwrap() {
Choice::Left => match workspace.active(*lhs) {
Some(new_lhs) => {
op = SsaOp::CopyReg(new_index, new_lhs);
}
None => {
workspace.set_active(*lhs, new_index);
continue;
}
},
Choice::Right => match workspace.active(*rhs) {
Some(new_rhs) => {
op = SsaOp::CopyReg(new_index, new_rhs);
}
None => {
workspace.set_active(*rhs, new_index);
continue;
}
},
Choice::Both => {
choice_count += 1;
*index = new_index;
*lhs = workspace.get_or_insert_active(*lhs);
*rhs = workspace.get_or_insert_active(*rhs);
}
Choice::Unknown => panic!("oh no"),
}
}
SsaOp::AddRegReg(index, lhs, rhs)
| SsaOp::MulRegReg(index, lhs, rhs)
| SsaOp::SubRegReg(index, lhs, rhs)
| SsaOp::DivRegReg(index, lhs, rhs)
| SsaOp::AtanRegReg(index, lhs, rhs)
| SsaOp::CompareRegReg(index, lhs, rhs)
| SsaOp::ModRegReg(index, lhs, rhs) => {
*index = new_index;
*lhs = workspace.get_or_insert_active(*lhs);
*rhs = workspace.get_or_insert_active(*rhs);
}
SsaOp::AddRegImm(index, arg, _imm)
| SsaOp::MulRegImm(index, arg, _imm)
| SsaOp::SubRegImm(index, arg, _imm)
| SsaOp::SubImmReg(index, arg, _imm)
| SsaOp::DivRegImm(index, arg, _imm)
| SsaOp::DivImmReg(index, arg, _imm)
| SsaOp::AtanImmReg(index, arg, _imm)
| SsaOp::AtanRegImm(index, arg, _imm)
| SsaOp::CompareRegImm(index, arg, _imm)
| SsaOp::CompareImmReg(index, arg, _imm)
| SsaOp::ModRegImm(index, arg, _imm)
| SsaOp::ModImmReg(index, arg, _imm) => {
*index = new_index;
*arg = workspace.get_or_insert_active(*arg);
}
}
workspace.alloc.op(op);
ops_out.push(op);
}
assert_eq!(workspace.count as usize + 1, ops_out.len());
let asm_tape = workspace.alloc.finalize();
Ok(VmData {
ssa: SsaTape {
tape: ops_out,
choice_count,
output_count,
},
asm: asm_tape,
vars: self.vars.clone(),
})
}
/// Produces an iterator that visits [`RegOp`] values in evaluation order
pub fn iter_asm(&self) -> impl Iterator<Item = RegOp> + '_ {
self.asm.iter().cloned().rev()
}
/// Pretty-prints the inner SSA tape
pub fn pretty_print(&self) {
self.ssa.pretty_print();
for a in self.iter_asm() {
println!("{a:?}");
}
}
}
////////////////////////////////////////////////////////////////////////////////
/// Data structures used during [`VmData::simplify`]
///
/// This is exposed to minimize reallocations in hot loops.
pub struct VmWorkspace<const N: usize> {
/// Register allocator
pub(crate) alloc: RegisterAllocator<N>,
/// Current bindings from SSA variables to registers
pub(crate) bind: Vec<u32>,
/// Number of active SSA bindings
///
/// This value is monotonically increasing; each SSA variable gets the next
/// value if it is unassigned when encountered.
count: u32,
}
impl<const N: usize> Default for VmWorkspace<N> {
fn default() -> Self {
Self {
alloc: RegisterAllocator::empty(),
bind: vec![],
count: 0,
}
}
}
impl<const N: usize> VmWorkspace<N> {
fn active(&self, i: u32) -> Option<u32> {
if self.bind[i as usize] != u32::MAX {
Some(self.bind[i as usize])
} else {
None
}
}
fn get_or_insert_active(&mut self, i: u32) -> u32 {
if self.bind[i as usize] == u32::MAX {
self.bind[i as usize] = self.count;
self.count += 1;
}
self.bind[i as usize]
}
fn set_active(&mut self, i: u32, bind: u32) {
self.bind[i as usize] = bind;
}
/// Resets the workspace, preserving allocations and claiming the given
/// [`RegTape`].
pub fn reset(&mut self, tape_len: usize, tape: RegTape) {
self.alloc.reset(tape_len, tape);
self.bind.fill(u32::MAX);
self.bind.resize(tape_len, u32::MAX);
self.count = 0;
}
}
#[cfg(test)]
mod test {
use super::*;
#[test]
fn simplify_reg_count_change() {
let mut ctx = Context::new();
let x = ctx.x();
let y = ctx.y();
let z = ctx.z();
let xy = ctx.add(x, y).unwrap();
let xyz = ctx.add(xy, z).unwrap();
let data = VmData::<3>::new(&ctx, &[xyz]).unwrap();
assert_eq!(data.len(), 6); // 3x input, 2x add, 1x output
let next = data
.simplify::<2>(&[], &mut Default::default(), Default::default())
.unwrap();
assert_eq!(next.len(), 8); // extra load + store
let data = VmData::<2>::new(&ctx, &[xyz]).unwrap();
assert_eq!(data.len(), 8);
let next = data
.simplify::<3>(&[], &mut Default::default(), Default::default())
.unwrap();
assert_eq!(next.len(), 6);
}
}
+2
View File
@@ -0,0 +1,2 @@
pub mod choice;
// mod data;