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
+17 -5
View File
@@ -41,16 +41,21 @@ for shader, ty in normal_shaders:
if result.returncode != 0: if result.returncode != 0:
sys.exit(result.returncode) sys.exit(result.returncode)
replacement_shaders = [("trace.frag", "frag"), ("normals.frag", "frag")] replacement_shaders = [
("trace.frag", "frag", (True, False, False)),
("normals.frag", "frag", (True, False, False)),
("fuzz.comp", "comp", (True, True, True)),
]
for shader, ty in replacement_shaders: for shader, ty, usage in replacement_shaders:
print(f"{SHADERS_IN}/{shader}.glsl") print(f"{SHADERS_IN}/{shader}.glsl")
(uses_point, uses_interval, uses_gradient) = usage
result = subprocess.run( result = subprocess.run(
[ [
"glslangValidator", "glslangValidator",
"--spirv-val", "--spirv-val",
"--spirv-dis", "--spirv-dis",
"-g0", "-gVS",
"-S", "-S",
ty, ty,
"--target-env", "--target-env",
@@ -62,12 +67,19 @@ for shader, ty in replacement_shaders:
text=True, text=True,
) )
if result.returncode != 0: if result.returncode != 0:
print(result.stdout)
sys.exit(result.returncode) sys.exit(result.returncode)
asm = result.stdout asm = result.stdout
with open(f"{SHADERS_OUT}/{shader}.old.spv-dis", "w") as f: with open(f"{SHADERS_OUT}/{shader}.old.spv-dis", "w") as f:
f.write(asm) f.write(asm)
asm = re.sub(r"(?ms)%11 =.*?OpFunctionEnd", "", asm)
asm = re.sub(r"%11 ", "%1000 ", asm) asm = re.sub(r"(?ms)%scene_vf4_ =.*?OpFunctionEnd", "", asm)
asm = re.sub(r"%scene_vf4_", "%10000", asm)
asm = re.sub(r"(?ms)%interval_scene_vf4_vf4_ =.*?OpFunctionEnd", "", asm)
asm = re.sub(r"%interval_scene_vf4_vf4_", "%20000", asm)
asm = re.sub(r"(?ms)%gradient_scene_vf4_ =.*?OpFunctionEnd", "", asm)
asm = re.sub(r"%gradient_scene_vf4_", "%30000", asm)
with open(f"{SHADERS_OUT}/{shader}.spv-dis", "w") as f: with open(f"{SHADERS_OUT}/{shader}.spv-dis", "w") as f:
f.write(asm) f.write(asm)
result = subprocess.run( result = subprocess.run(
+215 -1
View File
@@ -18,6 +18,7 @@ use std::{
fs::{File, remove_file, rename}, fs::{File, remove_file, rename},
io::{Cursor, Read, Write}, io::{Cursor, Read, Write},
path::{Path, PathBuf}, path::{Path, PathBuf},
simd::f32x8,
sync::{Arc, Mutex, atomic::AtomicBool, mpmc, mpsc}, sync::{Arc, Mutex, atomic::AtomicBool, mpmc, mpsc},
thread::{self, JoinHandle}, thread::{self, JoinHandle},
time::Instant, time::Instant,
@@ -106,7 +107,10 @@ use crate::{
trace_vs::{Camera, Lights, PushConstantData}, trace_vs::{Camera, Lights, PushConstantData},
}; };
mod objects; mod objects;
use tape_load::{interpreters, ssa, types}; use tape_load::{
interpreters::{self, interval::IntervalInterpreter, point::PointInterpreter},
ssa, types,
};
use crate::objects::*; use crate::objects::*;
@@ -217,6 +221,7 @@ struct App {
vertex_buffer: Subbuffer<[IVertex]>, vertex_buffer: Subbuffer<[IVertex]>,
trace_module: Module, trace_module: Module,
normals_module: Module, normals_module: Module,
fuzz_module: Module,
_threads: Vec<JoinHandle<()>>, _threads: Vec<JoinHandle<()>>,
thread_work_creation: mpmc::Sender<WorkItem>, thread_work_creation: mpmc::Sender<WorkItem>,
thread_work_completion: mpsc::Receiver<WorkComplete>, thread_work_completion: mpsc::Receiver<WorkComplete>,
@@ -598,6 +603,14 @@ impl App {
rspirv::binary::parse_words(normals_spv_code_u32, &mut loader).unwrap(); rspirv::binary::parse_words(normals_spv_code_u32, &mut loader).unwrap();
let normals_module = loader.module(); let normals_module = loader.module();
let fuzz_spv_code = include_bytes!("../shaders_out/fuzz.comp.spv");
let fuzz_spv_code_u32 = vulkano::shader::spirv::bytes_to_words(fuzz_spv_code)
.unwrap()
.into_owned();
let mut loader = rspirv::dr::Loader::new();
rspirv::binary::parse_words(fuzz_spv_code_u32, &mut loader).unwrap();
let fuzz_module = loader.module();
let (thread_work_creation_sender, thread_work_creation_receiver) = let (thread_work_creation_sender, thread_work_creation_receiver) =
mpmc::sync_channel::<WorkItem>(256); mpmc::sync_channel::<WorkItem>(256);
let (thread_work_completion_sender, thread_work_completion_receiver) = let (thread_work_completion_sender, thread_work_completion_receiver) =
@@ -668,6 +681,7 @@ impl App {
vertex_buffer, vertex_buffer,
trace_module, trace_module,
normals_module, normals_module,
fuzz_module,
_threads: threads, _threads: threads,
thread_work_creation: thread_work_creation_sender, thread_work_creation: thread_work_creation_sender,
thread_work_completion: thread_work_completion_receiver, thread_work_completion: thread_work_completion_receiver,
@@ -733,6 +747,15 @@ mod miss {
} }
} }
mod fuzz_cs {
vulkano_shaders::shader! {
bytes: "shaders_out/fuzz.comp.spv",
vulkan_version: "1.3",
spirv_version: "1.6",
custom_derives: [Debug, Clone, Copy],
}
}
impl ApplicationHandler for App { impl ApplicationHandler for App {
fn resumed(&mut self, event_loop: &ActiveEventLoop) { fn resumed(&mut self, event_loop: &ActiveEventLoop) {
let window = Arc::new( let window = Arc::new(
@@ -1157,6 +1180,7 @@ impl ApplicationHandler for App {
self.device.clone(), self.device.clone(),
self.trace_module.clone(), self.trace_module.clone(),
self.normals_module.clone(), self.normals_module.clone(),
self.fuzz_module.clone(),
rcx.render_pass.clone(), rcx.render_pass.clone(),
self.pipeline_cache.clone(), self.pipeline_cache.clone(),
rcx.shader_modules.clone(), rcx.shader_modules.clone(),
@@ -1687,12 +1711,22 @@ impl App {
} }
let miss = miss.unwrap(); let miss = miss.unwrap();
let fuzz_comp = read_spirv_words_from_file("fuzz.comp.spv");
if fuzz_comp.is_err() {
error!("Could not read fuzz comp file");
return;
}
let fuzz_comp = fuzz_comp.unwrap();
let mut trace_loader = rspirv::dr::Loader::new(); let mut trace_loader = rspirv::dr::Loader::new();
rspirv::binary::parse_words(trace_frag, &mut trace_loader).unwrap(); rspirv::binary::parse_words(trace_frag, &mut trace_loader).unwrap();
let mut normals_loader = rspirv::dr::Loader::new(); let mut normals_loader = rspirv::dr::Loader::new();
rspirv::binary::parse_words(normals_frag, &mut normals_loader).unwrap(); rspirv::binary::parse_words(normals_frag, &mut normals_loader).unwrap();
let mut fuzz_loader = rspirv::dr::Loader::new();
rspirv::binary::parse_words(fuzz_comp, &mut fuzz_loader).unwrap();
let trace_vert = unsafe { let trace_vert = unsafe {
::vulkano::shader::ShaderModule::new( ::vulkano::shader::ShaderModule::new(
self.device.clone(), self.device.clone(),
@@ -1752,6 +1786,7 @@ impl App {
self.trace_module = trace_loader.module(); self.trace_module = trace_loader.module();
self.normals_module = normals_loader.module(); self.normals_module = normals_loader.module();
self.fuzz_module = fuzz_loader.module();
} }
fn recreate_pipelines( fn recreate_pipelines(
@@ -1769,6 +1804,7 @@ impl App {
self.device.clone(), self.device.clone(),
self.trace_module.clone(), self.trace_module.clone(),
self.normals_module.clone(), self.normals_module.clone(),
self.fuzz_module.clone(),
rcx.render_pass.clone(), rcx.render_pass.clone(),
self.pipeline_cache.clone(), self.pipeline_cache.clone(),
rcx.shader_modules.clone(), rcx.shader_modules.clone(),
@@ -1800,6 +1836,180 @@ impl App {
} }
} }
fn compute_fuzz(&mut self, rcx: &mut RenderContext, work_for_later: &mut Vec<WorkComplete>) {
let mut test_csg = vec![];
let mut test_csg_count = 32;
for _ in 0..test_csg_count {
self.thread_work_creation
.send(WorkItem::CreateCSG(
self.device.clone(),
self.trace_module.clone(),
self.normals_module.clone(),
self.fuzz_module.clone(),
rcx.render_pass.clone(),
self.pipeline_cache.clone(),
rcx.shader_modules.clone(),
self.previous_debug,
self.block_enable_allocator.clone(),
self.host_visible_allocator.clone(),
self.descriptor_set_allocator.clone(),
rcx.swapchain.image_count(),
DEFAULT_SUBDIVISION,
rand::random(),
-1,
))
.unwrap();
}
const LOCAL_WIDTH: u32 = 128;
const GLOBAL_WIDTH: u32 = 32;
while test_csg_count > 0 {
for work in self.thread_work_completion.try_iter() {
match work {
WorkComplete::CreateCSG(ref csg, index) => {
if index < 0 {
let csg = csg.clone();
let alloc = self.host_visible_allocator.lock().unwrap();
let input: Subbuffer<[fuzz_cs::InputData]> = alloc
.allocate_slice((LOCAL_WIDTH * GLOBAL_WIDTH).into())
.unwrap();
let output: Subbuffer<[fuzz_cs::OutputData]> = alloc
.allocate_slice((LOCAL_WIDTH * GLOBAL_WIDTH).into())
.unwrap();
drop(alloc);
let mut input_rw = input.write().unwrap();
for val in input_rw.iter_mut() {
for i in 0..4 {
val.point[i] = rand::random_range(-1000.0..1000.0);
val.gradient[i] = rand::random_range(-1000.0..1000.0);
let midpoint = rand::random_range(-1000.0..1000.0);
val.interval_low[i] = rand::random_range(-1000.0..midpoint);
val.interval_high[i] = rand::random_range(midpoint..1000.0);
}
}
drop(input_rw);
let csg_rw = csg.read().unwrap();
let fuzz_layout =
csg_rw.fuzz_pipeline.layout().set_layouts()[2].clone();
drop(csg_rw);
let fuzz_descriptor_set = DescriptorSet::new(
self.descriptor_set_allocator.clone(),
fuzz_layout.clone(),
[
WriteDescriptorSet::buffer(0, input.clone()),
WriteDescriptorSet::buffer(1, output.clone()),
],
[],
)
.unwrap();
test_csg.push((csg, fuzz_descriptor_set, input, output));
test_csg_count -= 1;
} else {
work_for_later.push(work);
}
},
other => work_for_later.push(other),
}
}
}
let mut builder = AutoCommandBufferBuilder::primary(
self.command_buffer_allocator.clone(),
self.compute_queue.queue_family_index(),
CommandBufferUsage::OneTimeSubmit,
)
.unwrap();
for (csg, set, ..) in test_csg.iter() {
let csg = csg.read().unwrap();
builder
.bind_pipeline_compute(csg.fuzz_pipeline.clone())
.unwrap()
.bind_descriptor_sets(
PipelineBindPoint::Compute,
csg.fuzz_pipeline.layout().clone(),
2, // 2 is the index of our set
set.clone(),
)
.unwrap();
unsafe { builder.dispatch([GLOBAL_WIDTH, 1, 1]) }.unwrap();
}
let command_buffer = builder.build().unwrap();
let future = sync::now(self.device.clone())
.then_execute(self.compute_queue.clone(), command_buffer)
.unwrap()
.then_signal_fence_and_flush()
.unwrap();
future.wait(None).unwrap();
for (csg, _, input, output) in test_csg {
let csg = csg.read().unwrap();
let input = input.read().unwrap();
let output = output.read().unwrap();
for (input, output) in input
.as_chunks::<8>()
.0
.into_iter()
.zip(output.as_chunks::<8>().0)
{
println!("{output:?}");
let point_val = f32x8::from_slice(&output.map(|o| *o.point));
let point_expected = PointInterpreter::new(&csg.parts).scene(
f32x8::from_slice(&input.map(|i| i.point[0])),
f32x8::from_slice(&input.map(|i| i.point[1])),
f32x8::from_slice(&input.map(|i| i.point[2])),
f32x8::from_slice(&input.map(|i| i.point[3])),
);
if point_expected != point_val {
println!(
"ERROR: point expected {:?}, got {:?}",
point_expected, point_val
);
}
let interval_val = types::Interval::new(
f32x8::from_slice(&output.map(|o| o.interval[0])),
f32x8::from_slice(&output.map(|o| o.interval[1])),
);
let interval_expected = IntervalInterpreter::new(&csg.parts).scene(
types::Interval::new(
f32x8::from_slice(&input.map(|i| i.interval_low[0])),
f32x8::from_slice(&input.map(|i| i.interval_high[0])),
),
types::Interval::new(
f32x8::from_slice(&input.map(|i| i.interval_low[1])),
f32x8::from_slice(&input.map(|i| i.interval_high[1])),
),
types::Interval::new(
f32x8::from_slice(&input.map(|i| i.interval_low[2])),
f32x8::from_slice(&input.map(|i| i.interval_high[2])),
),
types::Interval::new(
f32x8::from_slice(&input.map(|i| i.interval_low[3])),
f32x8::from_slice(&input.map(|i| i.interval_high[3])),
),
);
if interval_expected != interval_val {
println!(
"ERROR: interval expected {:?}, got {:?}",
interval_expected, interval_val
);
}
}
}
}
fn redraw(&mut self, rcx: &mut RenderContext) { fn redraw(&mut self, rcx: &mut RenderContext) {
let window_size = rcx.window.inner_size(); let window_size = rcx.window.inner_size();
if window_size.width == 0 || window_size.height == 0 { if window_size.width == 0 || window_size.height == 0 {
@@ -1838,6 +2048,8 @@ impl App {
let mut work_for_later = vec![]; let mut work_for_later = vec![];
self.compute_fuzz(rcx, &mut work_for_later);
if self.draw_gui { if self.draw_gui {
gui_up(&mut rcx.gui, &mut self.gstate); gui_up(&mut rcx.gui, &mut self.gstate);
} }
@@ -1848,6 +2060,7 @@ impl App {
self.device.clone(), self.device.clone(),
self.trace_module.clone(), self.trace_module.clone(),
self.normals_module.clone(), self.normals_module.clone(),
self.fuzz_module.clone(),
rcx.render_pass.clone(), rcx.render_pass.clone(),
self.pipeline_cache.clone(), self.pipeline_cache.clone(),
rcx.shader_modules.clone(), rcx.shader_modules.clone(),
@@ -1938,6 +2151,7 @@ impl App {
self.device.clone(), self.device.clone(),
self.trace_module.clone(), self.trace_module.clone(),
self.normals_module.clone(), self.normals_module.clone(),
self.fuzz_module.clone(),
rcx.render_pass.clone(), rcx.render_pass.clone(),
self.pipeline_cache.clone(), self.pipeline_cache.clone(),
rcx.shader_modules.clone(), rcx.shader_modules.clone(),
+3 -1
View File
@@ -9,7 +9,7 @@ use vulkano::{
memory::allocator::{ memory::allocator::{
AllocationCreateInfo, MemoryAllocatePreference, MemoryAllocator, MemoryTypeFilter, AllocationCreateInfo, MemoryAllocatePreference, MemoryAllocator, MemoryTypeFilter,
}, },
pipeline::{GraphicsPipeline, graphics::vertex_input::Vertex}, pipeline::{ComputePipeline, GraphicsPipeline, graphics::vertex_input::Vertex},
shader::ShaderModule, shader::ShaderModule,
}; };
@@ -51,8 +51,10 @@ pub(crate) struct CSG {
pub(crate) new_pipelines_needed: bool, pub(crate) new_pipelines_needed: bool,
pub(crate) trace_shader_module: Arc<ShaderModule>, pub(crate) trace_shader_module: Arc<ShaderModule>,
pub(crate) normals_shader_module: Arc<ShaderModule>, pub(crate) normals_shader_module: Arc<ShaderModule>,
pub(crate) fuzz_shader_module: Arc<ShaderModule>,
pub(crate) trace_pipeline: Arc<GraphicsPipeline>, pub(crate) trace_pipeline: Arc<GraphicsPipeline>,
pub(crate) normals_pipeline: Arc<GraphicsPipeline>, pub(crate) normals_pipeline: Arc<GraphicsPipeline>,
pub(crate) fuzz_pipeline: Arc<ComputePipeline>,
pub(crate) subdivision: u32, pub(crate) subdivision: u32,
pub(crate) enable_buffer: Vec<Subbuffer<Object>>, pub(crate) enable_buffer: Vec<Subbuffer<Object>>,
pub(crate) enable_buffer_host_visible: Vec<Subbuffer<Object>>, pub(crate) enable_buffer_host_visible: Vec<Subbuffer<Object>>,
+36
View File
@@ -0,0 +1,36 @@
#version 460
#extension GL_GOOGLE_include_directive:require
#include "include.glsl"
#include "implicit_include.glsl"
layout(local_size_x = 128, local_size_y = 1, local_size_z = 1) in;
struct InputData {
vec4 point;
vec4 interval_low;
vec4 interval_high;
vec4 gradient;
};
struct OutputData {
float point;
vec2 interval;
vec4 gradient;
};
layout(set = 2, binding = 0, std430) restrict readonly buffer InputDataBuffer {
InputData data[];
} input_data;
layout(set = 2, binding = 1, std430) restrict writeonly buffer OutputDataBuffer {
OutputData data[];
} output_data;
void main() {
InputData i = input_data.data[gl_GlobalInvocationID.x];
OutputData o;
o.point = scene(i.point);
o.interval = interval_scene(i.interval_low, i.interval_high);
o.gradient = gradient_scene(i.gradient);
output_data.data[gl_GlobalInvocationID.x] = o;
}
@@ -24,4 +24,8 @@ vec2 interval_scene(vec4 pl, vec4 ph) {
return vec2(0.0); return vec2(0.0);
} }
vec4 gradient_scene(vec4 p) {
return vec4(0.0);
}
#endif #endif
+49 -6
View File
@@ -18,9 +18,10 @@ use vulkano::{
descriptor_set::{DescriptorSet, WriteDescriptorSet, allocator::DescriptorSetAllocator}, descriptor_set::{DescriptorSet, WriteDescriptorSet, allocator::DescriptorSetAllocator},
device::{Device, Queue}, device::{Device, Queue},
pipeline::{ pipeline::{
DynamicState, GraphicsPipeline, Pipeline, PipelineCreateFlags, PipelineLayout, ComputePipeline, DynamicState, GraphicsPipeline, Pipeline, PipelineCreateFlags,
PipelineShaderStageCreateInfo, PipelineLayout, PipelineShaderStageCreateInfo,
cache::PipelineCache, cache::PipelineCache,
compute::ComputePipelineCreateInfo,
graphics::{ graphics::{
GraphicsPipelineCreateInfo, GraphicsPipelineCreateInfo,
color_blend::{ColorBlendAttachmentState, ColorBlendState}, color_blend::{ColorBlendAttachmentState, ColorBlendState},
@@ -54,6 +55,7 @@ pub enum WorkItem {
Arc<Device>, Arc<Device>,
Module, Module,
Module, Module,
Module,
Arc<RenderPass>, Arc<RenderPass>,
Arc<PipelineCache>, Arc<PipelineCache>,
ShaderModules, ShaderModules,
@@ -79,6 +81,7 @@ pub enum WorkItem {
Arc<Device>, Arc<Device>,
Module, Module,
Module, Module,
Module,
Arc<RenderPass>, Arc<RenderPass>,
Arc<PipelineCache>, Arc<PipelineCache>,
ShaderModules, ShaderModules,
@@ -100,6 +103,7 @@ pub fn thread_loop(recv: mpmc::Receiver<WorkItem>, send: mpsc::SyncSender<WorkCo
device, device,
trace_module, trace_module,
normals_module, normals_module,
fuzz_module,
render_pass, render_pass,
cache, cache,
modules, modules,
@@ -119,14 +123,18 @@ pub fn thread_loop(recv: mpmc::Receiver<WorkItem>, send: mpsc::SyncSender<WorkCo
sdf_specialize_module(device.clone(), &parts, trace_module, "trace"); sdf_specialize_module(device.clone(), &parts, trace_module, "trace");
let normals_shader_module = let normals_shader_module =
sdf_specialize_module(device.clone(), &parts, normals_module, "normals"); sdf_specialize_module(device.clone(), &parts, normals_module, "normals");
let fuzz_shader_module =
sdf_specialize_module(device.clone(), &parts, fuzz_module, "fuzz");
let (trace_pipeline, normals_pipeline) = deferred_pipelines_recompile( let (trace_pipeline, normals_pipeline, fuzz_pipeline) =
deferred_pipelines_recompile(
device, device,
render_pass, render_pass,
cache, cache,
modules, modules,
trace_shader_module.clone(), trace_shader_module.clone(),
normals_shader_module.clone(), normals_shader_module.clone(),
fuzz_shader_module.clone(),
debug, debug,
); );
@@ -191,8 +199,10 @@ pub fn thread_loop(recv: mpmc::Receiver<WorkItem>, send: mpsc::SyncSender<WorkCo
new_pipelines_needed: false, new_pipelines_needed: false,
trace_shader_module, trace_shader_module,
normals_shader_module, normals_shader_module,
fuzz_shader_module,
trace_pipeline, trace_pipeline,
normals_pipeline, normals_pipeline,
fuzz_pipeline,
subdivision, subdivision,
enable_buffer, enable_buffer,
enable_buffer_host_visible, enable_buffer_host_visible,
@@ -256,6 +266,7 @@ pub fn thread_loop(recv: mpmc::Receiver<WorkItem>, send: mpsc::SyncSender<WorkCo
device, device,
trace_module, trace_module,
normals_module, normals_module,
fuzz_module,
render_pass, render_pass,
cache, cache,
modules, modules,
@@ -268,14 +279,18 @@ pub fn thread_loop(recv: mpmc::Receiver<WorkItem>, send: mpsc::SyncSender<WorkCo
sdf_specialize_module(device.clone(), &csg.parts, trace_module, "trace"); sdf_specialize_module(device.clone(), &csg.parts, trace_module, "trace");
csg.normals_shader_module = csg.normals_shader_module =
sdf_specialize_module(device.clone(), &csg.parts, normals_module, "normals"); sdf_specialize_module(device.clone(), &csg.parts, normals_module, "normals");
csg.fuzz_shader_module =
sdf_specialize_module(device.clone(), &csg.parts, fuzz_module, "fuzz");
(csg.trace_pipeline, csg.normals_pipeline) = deferred_pipelines_recompile( (csg.trace_pipeline, csg.normals_pipeline, csg.fuzz_pipeline) =
deferred_pipelines_recompile(
device, device,
render_pass, render_pass,
cache, cache,
modules, modules,
csg.trace_shader_module.clone(), csg.trace_shader_module.clone(),
csg.normals_shader_module.clone(), csg.normals_shader_module.clone(),
csg.fuzz_shader_module.clone(),
debug, debug,
); );
csg.new_pipelines_needed = false; csg.new_pipelines_needed = false;
@@ -298,8 +313,13 @@ fn deferred_pipelines_recompile(
shader_modules: ShaderModules, shader_modules: ShaderModules,
trace_shader_module: Arc<ShaderModule>, trace_shader_module: Arc<ShaderModule>,
normals_shader_module: Arc<ShaderModule>, normals_shader_module: Arc<ShaderModule>,
fuzz_shader_module: Arc<ShaderModule>,
debug: PreviousDebug, debug: PreviousDebug,
) -> (Arc<GraphicsPipeline>, Arc<GraphicsPipeline>) { ) -> (
Arc<GraphicsPipeline>,
Arc<GraphicsPipeline>,
Arc<ComputePipeline>,
) {
let specs = get_spec_constants(&debug); let specs = get_spec_constants(&debug);
let dynamic_state = [DynamicState::Viewport].into_iter().collect::<HashSet<_>>(); let dynamic_state = [DynamicState::Viewport].into_iter().collect::<HashSet<_>>();
@@ -426,7 +446,30 @@ fn deferred_pipelines_recompile(
) )
.unwrap(); .unwrap();
(trace_pipeline, normals_pipeline) let fuzz_cs_info = PipelineShaderStageCreateInfo::new(
fuzz_shader_module
.specialize(specs.clone())
.unwrap()
.single_entry_point()
.unwrap(),
);
let compute_pipeline_layout = PipelineLayout::new(
device.clone(),
PipelineDescriptorSetLayoutCreateInfo::from_stages([&fuzz_cs_info])
.into_pipeline_layout_create_info(device.clone())
.unwrap(),
)
.unwrap();
let compute_pipeline = ComputePipeline::new(
device.clone(),
None,
ComputePipelineCreateInfo::stage_layout(fuzz_cs_info, compute_pipeline_layout.clone()),
)
.expect("failed to create compute pipeline");
(trace_pipeline, normals_pipeline, compute_pipeline)
} }
fn create_csg(seed: u64) -> SSATape { fn create_csg(seed: u64) -> SSATape {
+10 -34
View File
@@ -20,7 +20,7 @@ pub(crate) fn compile_gradient_function(
spirv::FunctionControl::INLINE spirv::FunctionControl::INLINE
| spirv::FunctionControl::PURE | spirv::FunctionControl::PURE
| spirv::FunctionControl::CONST, | spirv::FunctionControl::CONST,
types.point_fn_type, types.gradient_fn_type,
) )
.unwrap(); .unwrap();
let pos_p = b.function_parameter(types.vec4p).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 { match instruction.opcode.opcode {
SSAStop => { SSAStop => {
let zero = b.constant_bit32(types.float, (0.0f32).to_bits()); 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 => { SSAReturn => {
let value = input_resolve(types.float, b, &mapping, instruction.inputs[0]); 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 => { SSAPosition => {
mapping.insert( 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 interval;
pub(crate) mod point; pub(crate) mod point;
@@ -15,5 +15,6 @@ pub(crate) struct SpirVTypes {
pub vec4p: u32, pub vec4p: u32,
pub point_fn_type: u32, pub point_fn_type: u32,
pub interval_fn_type: u32, pub interval_fn_type: u32,
pub gradient_fn_type: u32,
pub jit_string: u32, pub jit_string: u32,
} }
+23 -16
View File
@@ -6,7 +6,8 @@ use crate::{
DUMP_SPV_DIS_TO_FILE, DUMP_SPV_DIS_TO_FILE,
instruction_set::InstructionSet, instruction_set::InstructionSet,
spirv_compilers::{ 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 vec4p = b.type_pointer(None, spirv::StorageClass::Function, vec4);
let point_fn_type = b.type_function(float, vec![vec4p]); let point_fn_type = b.type_function(float, vec![vec4p]);
let interval_fn_type = b.type_function(vec2, vec![vec4p, 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"); let jit_string = b.string("JIT");
const POINT_FN_LOC: u32 = 1000; const POINT_FN_LOC: u32 = 10000;
const INTERVAL_FN_LOC: u32 = 2000; const INTERVAL_FN_LOC: u32 = 20000;
const GRADIENT_FN_LOC: u32 = 3000; const GRADIENT_FN_LOC: u32 = 30000;
let types = SpirVTypes { let types = SpirVTypes {
glsl, glsl,
@@ -499,6 +501,7 @@ impl SSATape {
vec4p, vec4p,
point_fn_type, point_fn_type,
interval_fn_type, interval_fn_type,
gradient_fn_type,
jit_string, jit_string,
}; };
@@ -536,18 +539,22 @@ impl SSATape {
}, },
); );
//// Manually fix the header bounds // Manually fix the header bounds
//let mut module = b.module(); let mut module = b.module();
//let header = module.header.as_mut().unwrap(); let header = module.header.as_mut().unwrap();
//header.bound = header.bound.max(GRADIENT_FN_LOC + 1); header.bound = header.bound.max(GRADIENT_FN_LOC + 1);
//let mut b = rspirv::dr::Builder::new_from_module(module); let mut b = rspirv::dr::Builder::new_from_module(module);
//
//compile_gradient_function( compile_gradient_function(
// &mut b, &mut b,
// self, self,
// types, types,
// if with_module { Some(GRADIENT_FN_LOC) } else { None }, if with_module {
//); Some(GRADIENT_FN_LOC)
} else {
None
},
);
let module = b.module(); let module = b.module();