compute fuxxing?

This commit is contained in:
2025-12-27 16:41:32 +00:00
parent 1c897fbfb7
commit 589b8c74d9
9 changed files with 375 additions and 80 deletions
+10 -34
View File
@@ -20,7 +20,7 @@ pub(crate) fn compile_gradient_function(
spirv::FunctionControl::INLINE
| spirv::FunctionControl::PURE
| spirv::FunctionControl::CONST,
types.point_fn_type,
types.gradient_fn_type,
)
.unwrap();
let pos_p = b.function_parameter(types.vec4p).unwrap();
@@ -106,45 +106,21 @@ pub(crate) fn compile_gradient_function(
}
}
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();
let out = b
.composite_construct(types.vec4, None, vec![zero, zero, zero, zero])
.unwrap();
b.ret_value(out).unwrap();
},
SSAReturn => {
let value = input_resolve(types.float, b, &mapping, instruction.inputs[0]);
b.ret_value(value).unwrap();
let zero = b.constant_bit32(types.float, (0.0f32).to_bits());
let out = b
.composite_construct(types.vec4, None, vec![value, zero, zero, zero])
.unwrap();
b.ret_value(out).unwrap();
},
SSAPosition => {
mapping.insert(
+2 -1
View File
@@ -1,4 +1,4 @@
//pub(crate) mod gradient;
pub(crate) mod gradient;
pub(crate) mod interval;
pub(crate) mod point;
@@ -15,5 +15,6 @@ pub(crate) struct SpirVTypes {
pub vec4p: u32,
pub point_fn_type: u32,
pub interval_fn_type: u32,
pub gradient_fn_type: u32,
pub jit_string: u32,
}
+23 -16
View File
@@ -6,7 +6,8 @@ use crate::{
DUMP_SPV_DIS_TO_FILE,
instruction_set::InstructionSet,
spirv_compilers::{
SpirVTypes, interval::compile_interval_function, point::compile_point_function,
SpirVTypes, gradient::compile_gradient_function, interval::compile_interval_function,
point::compile_point_function,
},
};
@@ -481,11 +482,12 @@ impl SSATape {
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 gradient_fn_type = b.type_function(vec4, vec![vec4p]);
let jit_string = b.string("JIT");
const POINT_FN_LOC: u32 = 1000;
const INTERVAL_FN_LOC: u32 = 2000;
const GRADIENT_FN_LOC: u32 = 3000;
const POINT_FN_LOC: u32 = 10000;
const INTERVAL_FN_LOC: u32 = 20000;
const GRADIENT_FN_LOC: u32 = 30000;
let types = SpirVTypes {
glsl,
@@ -499,6 +501,7 @@ impl SSATape {
vec4p,
point_fn_type,
interval_fn_type,
gradient_fn_type,
jit_string,
};
@@ -536,18 +539,22 @@ impl SSATape {
},
);
//// Manually fix the header bounds
//let mut module = b.module();
//let header = module.header.as_mut().unwrap();
//header.bound = header.bound.max(GRADIENT_FN_LOC + 1);
//let mut b = rspirv::dr::Builder::new_from_module(module);
//
//compile_gradient_function(
// &mut b,
// self,
// types,
// if with_module { Some(GRADIENT_FN_LOC) } else { None },
//);
// Manually fix the header bounds
let mut module = b.module();
let header = module.header.as_mut().unwrap();
header.bound = header.bound.max(GRADIENT_FN_LOC + 1);
let mut b = rspirv::dr::Builder::new_from_module(module);
compile_gradient_function(
&mut b,
self,
types,
if with_module {
Some(GRADIENT_FN_LOC)
} else {
None
},
);
let module = b.module();