Fix compute fuzzing

This commit is contained in:
2025-12-27 17:54:40 +00:00
parent 589b8c74d9
commit 6e6d74ef5a
7 changed files with 67 additions and 44 deletions
+11 -5
View File
@@ -10,6 +10,8 @@ SHADERS_OUT = "shaders_out"
SHADERS_IN = "src/shaders" SHADERS_IN = "src/shaders"
TARGET_ENV = "vulkan1.3" TARGET_ENV = "vulkan1.3"
DUMP_DISASSEMBLIES = True
if os.path.isdir(SHADERS_OUT): if os.path.isdir(SHADERS_OUT):
shutil.rmtree(SHADERS_OUT) shutil.rmtree(SHADERS_OUT)
os.mkdir(SHADERS_OUT) os.mkdir(SHADERS_OUT)
@@ -55,7 +57,7 @@ for shader, ty, usage in replacement_shaders:
"glslangValidator", "glslangValidator",
"--spirv-val", "--spirv-val",
"--spirv-dis", "--spirv-dis",
"-gVS", "-g",
"-S", "-S",
ty, ty,
"--target-env", "--target-env",
@@ -70,8 +72,10 @@ for shader, ty, usage in replacement_shaders:
print(result.stdout) 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:
f.write(asm) if DUMP_DISASSEMBLIES:
with open(f"{SHADERS_OUT}/{shader}.stage0.spv-dis", "w") as f:
f.write(asm)
asm = re.sub(r"(?ms)%scene_vf4_ =.*?OpFunctionEnd", "", asm) asm = re.sub(r"(?ms)%scene_vf4_ =.*?OpFunctionEnd", "", asm)
asm = re.sub(r"%scene_vf4_", "%10000", asm) asm = re.sub(r"%scene_vf4_", "%10000", asm)
@@ -80,8 +84,10 @@ for shader, ty, usage in replacement_shaders:
asm = re.sub(r"(?ms)%gradient_scene_vf4_ =.*?OpFunctionEnd", "", asm) asm = re.sub(r"(?ms)%gradient_scene_vf4_ =.*?OpFunctionEnd", "", asm)
asm = re.sub(r"%gradient_scene_vf4_", "%30000", asm) asm = re.sub(r"%gradient_scene_vf4_", "%30000", asm)
with open(f"{SHADERS_OUT}/{shader}.spv-dis", "w") as f: if DUMP_DISASSEMBLIES:
f.write(asm) with open(f"{SHADERS_OUT}/{shader}.stage1.spv-dis", "w") as f:
f.write(asm)
result = subprocess.run( result = subprocess.run(
[ [
"spirv-as", "spirv-as",
+32 -9
View File
@@ -5,6 +5,7 @@ const DUMP_SPV_TO_FILE: bool = false;
const PIPELINE_CACHING: bool = false; const PIPELINE_CACHING: bool = false;
const IGNORE_STENCIL: bool = true; const IGNORE_STENCIL: bool = true;
const VSYNC: bool = true; const VSYNC: bool = true;
const COMPUTE_FUZZING: bool = false;
const DEFAULT_SUBDIVISION: u32 = 16; const DEFAULT_SUBDIVISION: u32 = 16;
const MAXIMUM_SUBDIVISION: u32 = 64; const MAXIMUM_SUBDIVISION: u32 = 64;
@@ -18,7 +19,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, simd::{f32x8, num::SimdFloat},
sync::{Arc, Mutex, atomic::AtomicBool, mpmc, mpsc}, sync::{Arc, Mutex, atomic::AtomicBool, mpmc, mpsc},
thread::{self, JoinHandle}, thread::{self, JoinHandle},
time::Instant, time::Instant,
@@ -211,6 +212,7 @@ struct App {
uniform_buffer_allocator: Arc<Mutex<SubbufferAllocator>>, uniform_buffer_allocator: Arc<Mutex<SubbufferAllocator>>,
block_enable_allocator: Arc<Mutex<SubbufferAllocator>>, block_enable_allocator: Arc<Mutex<SubbufferAllocator>>,
host_visible_allocator: Arc<Mutex<SubbufferAllocator>>, host_visible_allocator: Arc<Mutex<SubbufferAllocator>>,
compute_visible_allocator: Arc<Mutex<SubbufferAllocator>>,
pipeline_cache: Arc<PipelineCache>, pipeline_cache: Arc<PipelineCache>,
draw_gui: bool, draw_gui: bool,
gstate: GState, gstate: GState,
@@ -538,6 +540,16 @@ impl App {
}, },
))); )));
let compute_visible_allocator = Arc::new(Mutex::new(SubbufferAllocator::new(
memory_allocator.clone(),
SubbufferAllocatorCreateInfo {
buffer_usage: BufferUsage::STORAGE_BUFFER,
memory_type_filter: MemoryTypeFilter::HOST_SEQUENTIAL_WRITE
| MemoryTypeFilter::PREFER_DEVICE,
..Default::default()
},
)));
let memory_allocator: Arc<dyn MemoryAllocator> = memory_allocator; let memory_allocator: Arc<dyn MemoryAllocator> = memory_allocator;
let pipeline_cache = get_pipeline_cache(device.clone()); let pipeline_cache = get_pipeline_cache(device.clone());
@@ -669,6 +681,7 @@ impl App {
descriptor_set_allocator, descriptor_set_allocator,
command_buffer_allocator, command_buffer_allocator,
uniform_buffer_allocator, uniform_buffer_allocator,
compute_visible_allocator,
host_visible_allocator, host_visible_allocator,
block_enable_allocator, block_enable_allocator,
pipeline_cache, pipeline_cache,
@@ -1871,7 +1884,7 @@ impl App {
if index < 0 { if index < 0 {
let csg = csg.clone(); let csg = csg.clone();
let alloc = self.host_visible_allocator.lock().unwrap(); let alloc = self.compute_visible_allocator.lock().unwrap();
let input: Subbuffer<[fuzz_cs::InputData]> = alloc let input: Subbuffer<[fuzz_cs::InputData]> = alloc
.allocate_slice((LOCAL_WIDTH * GLOBAL_WIDTH).into()) .allocate_slice((LOCAL_WIDTH * GLOBAL_WIDTH).into())
.unwrap(); .unwrap();
@@ -1885,7 +1898,7 @@ impl App {
for i in 0..4 { for i in 0..4 {
val.point[i] = rand::random_range(-1000.0..1000.0); val.point[i] = rand::random_range(-1000.0..1000.0);
val.gradient[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); let midpoint = rand::random_range(-999.0..999.0);
val.interval_low[i] = rand::random_range(-1000.0..midpoint); val.interval_low[i] = rand::random_range(-1000.0..midpoint);
val.interval_high[i] = rand::random_range(midpoint..1000.0); val.interval_high[i] = rand::random_range(midpoint..1000.0);
} }
@@ -1961,7 +1974,7 @@ impl App {
.into_iter() .into_iter()
.zip(output.as_chunks::<8>().0) .zip(output.as_chunks::<8>().0)
{ {
println!("{output:?}"); //println!("{output:?}");
let point_val = f32x8::from_slice(&output.map(|o| *o.point)); let point_val = f32x8::from_slice(&output.map(|o| *o.point));
let point_expected = PointInterpreter::new(&csg.parts).scene( 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[0])),
@@ -1970,7 +1983,7 @@ impl App {
f32x8::from_slice(&input.map(|i| i.point[3])), f32x8::from_slice(&input.map(|i| i.point[3])),
); );
if point_expected != point_val { if (point_expected - point_val).abs() > f32x8::splat(0.001) {
println!( println!(
"ERROR: point expected {:?}, got {:?}", "ERROR: point expected {:?}, got {:?}",
point_expected, point_val point_expected, point_val
@@ -2000,10 +2013,18 @@ impl App {
), ),
); );
if interval_expected != interval_val { if (interval_expected.upper() - interval_val.upper()).abs() > f32x8::splat(0.001) {
println!( println!(
"ERROR: interval expected {:?}, got {:?}", "ERROR: interval upper expected {:?}, got {:?}",
interval_expected, interval_val interval_expected.upper(),
interval_val.upper()
);
}
if (interval_expected.lower() - interval_val.lower()).abs() > f32x8::splat(0.001) {
println!(
"ERROR: interval lower expected {:?}, got {:?}",
interval_expected.lower(),
interval_val.lower()
); );
} }
} }
@@ -2048,7 +2069,9 @@ impl App {
let mut work_for_later = vec![]; let mut work_for_later = vec![];
self.compute_fuzz(rcx, &mut work_for_later); if COMPUTE_FUZZING {
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);
+5 -3
View File
@@ -12,15 +12,17 @@ pub(crate) fn compile_gradient_function(
types: SpirVTypes, types: SpirVTypes,
function_id: Option<spirv::Word>, function_id: Option<spirv::Word>,
) { ) {
let gradient_fn_type = b.type_function(types.vec4, vec![types.vec4p]);
let jit_string = b.string("Gradient JIT");
let _scene = b let _scene = b
.begin_function( .begin_function(
types.float, types.vec4,
function_id, function_id,
//spirv::FunctionControl::DONT_INLINE //spirv::FunctionControl::DONT_INLINE
spirv::FunctionControl::INLINE spirv::FunctionControl::INLINE
| spirv::FunctionControl::PURE | spirv::FunctionControl::PURE
| spirv::FunctionControl::CONST, | spirv::FunctionControl::CONST,
types.gradient_fn_type, gradient_fn_type,
) )
.unwrap(); .unwrap();
let pos_p = b.function_parameter(types.vec4p).unwrap(); let pos_p = b.function_parameter(types.vec4p).unwrap();
@@ -35,7 +37,7 @@ pub(crate) fn compile_gradient_function(
use SSAOpcode::*; use SSAOpcode::*;
use rspirv::dr::Operand::IdRef; use rspirv::dr::Operand::IdRef;
b.line(types.jit_string, line as u32, 0); b.line(jit_string, line as u32, 0);
fn input_resolve( fn input_resolve(
float: u32, float: u32,
+4 -2
View File
@@ -13,6 +13,8 @@ pub(crate) fn compile_interval_function(
types: SpirVTypes, types: SpirVTypes,
function_id: Option<spirv::Word>, function_id: Option<spirv::Word>,
) { ) {
let interval_fn_type = b.type_function(types.vec2, vec![types.vec4p, types.vec4p]);
let jit_string = b.string("Interval JIT");
let _scene = b let _scene = b
.begin_function( .begin_function(
types.vec2, types.vec2,
@@ -21,7 +23,7 @@ pub(crate) fn compile_interval_function(
spirv::FunctionControl::INLINE spirv::FunctionControl::INLINE
| spirv::FunctionControl::PURE | spirv::FunctionControl::PURE
| spirv::FunctionControl::CONST, | spirv::FunctionControl::CONST,
types.interval_fn_type, interval_fn_type,
) )
.unwrap(); .unwrap();
let pos_low = b.function_parameter(types.vec4p).unwrap(); let pos_low = b.function_parameter(types.vec4p).unwrap();
@@ -38,7 +40,7 @@ pub(crate) fn compile_interval_function(
use SSAOpcode::*; use SSAOpcode::*;
use rspirv::dr::Operand::IdRef; use rspirv::dr::Operand::IdRef;
b.line(types.jit_string, line as u32, 0); b.line(jit_string, line as u32, 0);
fn input_resolve( fn input_resolve(
float: u32, float: u32,
+9 -13
View File
@@ -4,17 +4,13 @@ pub(crate) mod point;
#[derive(Debug, Clone, Copy)] #[derive(Debug, Clone, Copy)]
pub(crate) struct SpirVTypes { pub(crate) struct SpirVTypes {
pub glsl: u32, pub glsl: u32,
pub void: u32, pub void: u32,
pub float: u32, pub float: u32,
pub bool: u32, pub bool: u32,
pub choice: u32, pub choice: u32,
pub vec2: u32, pub vec2: u32,
pub vec3: u32, pub vec3: u32,
pub vec4: u32, pub vec4: u32,
pub vec4p: u32, pub vec4p: u32,
pub point_fn_type: u32,
pub interval_fn_type: u32,
pub gradient_fn_type: u32,
pub jit_string: u32,
} }
+4 -2
View File
@@ -12,6 +12,8 @@ pub(crate) fn compile_point_function(
types: SpirVTypes, types: SpirVTypes,
function_id: Option<spirv::Word>, function_id: Option<spirv::Word>,
) { ) {
let point_fn_type = b.type_function(types.float, vec![types.vec4p]);
let jit_string = b.string("Point JIT");
let _scene = b let _scene = b
.begin_function( .begin_function(
types.float, types.float,
@@ -20,7 +22,7 @@ pub(crate) fn compile_point_function(
spirv::FunctionControl::INLINE spirv::FunctionControl::INLINE
| spirv::FunctionControl::PURE | spirv::FunctionControl::PURE
| spirv::FunctionControl::CONST, | spirv::FunctionControl::CONST,
types.point_fn_type, point_fn_type,
) )
.unwrap(); .unwrap();
let pos_p = b.function_parameter(types.vec4p).unwrap(); let pos_p = b.function_parameter(types.vec4p).unwrap();
@@ -35,7 +37,7 @@ pub(crate) fn compile_point_function(
use SSAOpcode::*; use SSAOpcode::*;
use rspirv::dr::Operand::IdRef; use rspirv::dr::Operand::IdRef;
b.line(types.jit_string, line as u32, 0); b.line(jit_string, line as u32, 0);
fn input_resolve( fn input_resolve(
float: u32, float: u32,
+2 -10
View File
@@ -467,7 +467,7 @@ impl SSATape {
b b
}; };
let glsl = if with_module { let glsl = if with_module {
1 4
} else { } else {
b.ext_inst_import("GLSL.std.450") b.ext_inst_import("GLSL.std.450")
}; };
@@ -480,10 +480,6 @@ impl SSATape {
let vec3 = b.type_vector(float, 3); let vec3 = b.type_vector(float, 3);
let vec4 = b.type_vector(float, 4); let vec4 = b.type_vector(float, 4);
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 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 = 10000; const POINT_FN_LOC: u32 = 10000;
const INTERVAL_FN_LOC: u32 = 20000; const INTERVAL_FN_LOC: u32 = 20000;
@@ -499,10 +495,6 @@ impl SSATape {
vec3, vec3,
vec4, vec4,
vec4p, vec4p,
point_fn_type,
interval_fn_type,
gradient_fn_type,
jit_string,
}; };
// Manually fix the header bounds // Manually fix the header bounds
@@ -561,7 +553,7 @@ impl SSATape {
if DUMP_SPV_DIS_TO_FILE { if DUMP_SPV_DIS_TO_FILE {
std::fs::write( std::fs::write(
format!( format!(
"{}.spv-dis", "spv-dis/{}.spv-dis",
humantime::format_rfc3339(std::time::SystemTime::now()) humantime::format_rfc3339(std::time::SystemTime::now())
), ),
module.disassemble(), module.disassemble(),