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"
TARGET_ENV = "vulkan1.3"
DUMP_DISASSEMBLIES = True
if os.path.isdir(SHADERS_OUT):
shutil.rmtree(SHADERS_OUT)
os.mkdir(SHADERS_OUT)
@@ -55,7 +57,7 @@ for shader, ty, usage in replacement_shaders:
"glslangValidator",
"--spirv-val",
"--spirv-dis",
"-gVS",
"-g",
"-S",
ty,
"--target-env",
@@ -70,8 +72,10 @@ for shader, ty, usage in replacement_shaders:
print(result.stdout)
sys.exit(result.returncode)
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"%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"%gradient_scene_vf4_", "%30000", asm)
with open(f"{SHADERS_OUT}/{shader}.spv-dis", "w") as f:
f.write(asm)
if DUMP_DISASSEMBLIES:
with open(f"{SHADERS_OUT}/{shader}.stage1.spv-dis", "w") as f:
f.write(asm)
result = subprocess.run(
[
"spirv-as",
+32 -9
View File
@@ -5,6 +5,7 @@ const DUMP_SPV_TO_FILE: bool = false;
const PIPELINE_CACHING: bool = false;
const IGNORE_STENCIL: bool = true;
const VSYNC: bool = true;
const COMPUTE_FUZZING: bool = false;
const DEFAULT_SUBDIVISION: u32 = 16;
const MAXIMUM_SUBDIVISION: u32 = 64;
@@ -18,7 +19,7 @@ use std::{
fs::{File, remove_file, rename},
io::{Cursor, Read, Write},
path::{Path, PathBuf},
simd::f32x8,
simd::{f32x8, num::SimdFloat},
sync::{Arc, Mutex, atomic::AtomicBool, mpmc, mpsc},
thread::{self, JoinHandle},
time::Instant,
@@ -211,6 +212,7 @@ struct App {
uniform_buffer_allocator: Arc<Mutex<SubbufferAllocator>>,
block_enable_allocator: Arc<Mutex<SubbufferAllocator>>,
host_visible_allocator: Arc<Mutex<SubbufferAllocator>>,
compute_visible_allocator: Arc<Mutex<SubbufferAllocator>>,
pipeline_cache: Arc<PipelineCache>,
draw_gui: bool,
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 pipeline_cache = get_pipeline_cache(device.clone());
@@ -669,6 +681,7 @@ impl App {
descriptor_set_allocator,
command_buffer_allocator,
uniform_buffer_allocator,
compute_visible_allocator,
host_visible_allocator,
block_enable_allocator,
pipeline_cache,
@@ -1871,7 +1884,7 @@ impl App {
if index < 0 {
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
.allocate_slice((LOCAL_WIDTH * GLOBAL_WIDTH).into())
.unwrap();
@@ -1885,7 +1898,7 @@ impl App {
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);
let midpoint = rand::random_range(-999.0..999.0);
val.interval_low[i] = rand::random_range(-1000.0..midpoint);
val.interval_high[i] = rand::random_range(midpoint..1000.0);
}
@@ -1961,7 +1974,7 @@ impl App {
.into_iter()
.zip(output.as_chunks::<8>().0)
{
println!("{output:?}");
//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])),
@@ -1970,7 +1983,7 @@ impl App {
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!(
"ERROR: point expected {:?}, got {:?}",
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!(
"ERROR: interval expected {:?}, got {:?}",
interval_expected, interval_val
"ERROR: interval upper expected {:?}, got {:?}",
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![];
self.compute_fuzz(rcx, &mut work_for_later);
if COMPUTE_FUZZING {
self.compute_fuzz(rcx, &mut work_for_later);
}
if self.draw_gui {
gui_up(&mut rcx.gui, &mut self.gstate);
+5 -3
View File
@@ -12,15 +12,17 @@ pub(crate) fn compile_gradient_function(
types: SpirVTypes,
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
.begin_function(
types.float,
types.vec4,
function_id,
//spirv::FunctionControl::DONT_INLINE
spirv::FunctionControl::INLINE
| spirv::FunctionControl::PURE
| spirv::FunctionControl::CONST,
types.gradient_fn_type,
gradient_fn_type,
)
.unwrap();
let pos_p = b.function_parameter(types.vec4p).unwrap();
@@ -35,7 +37,7 @@ pub(crate) fn compile_gradient_function(
use SSAOpcode::*;
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(
float: u32,
+4 -2
View File
@@ -13,6 +13,8 @@ pub(crate) fn compile_interval_function(
types: SpirVTypes,
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
.begin_function(
types.vec2,
@@ -21,7 +23,7 @@ pub(crate) fn compile_interval_function(
spirv::FunctionControl::INLINE
| spirv::FunctionControl::PURE
| spirv::FunctionControl::CONST,
types.interval_fn_type,
interval_fn_type,
)
.unwrap();
let pos_low = b.function_parameter(types.vec4p).unwrap();
@@ -38,7 +40,7 @@ pub(crate) fn compile_interval_function(
use SSAOpcode::*;
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(
float: u32,
+9 -13
View File
@@ -4,17 +4,13 @@ pub(crate) mod point;
#[derive(Debug, Clone, Copy)]
pub(crate) struct SpirVTypes {
pub glsl: u32,
pub void: u32,
pub float: u32,
pub bool: u32,
pub choice: u32,
pub vec2: u32,
pub vec3: u32,
pub vec4: u32,
pub vec4p: u32,
pub point_fn_type: u32,
pub interval_fn_type: u32,
pub gradient_fn_type: u32,
pub jit_string: u32,
pub glsl: u32,
pub void: u32,
pub float: u32,
pub bool: u32,
pub choice: u32,
pub vec2: u32,
pub vec3: u32,
pub vec4: u32,
pub vec4p: u32,
}
+4 -2
View File
@@ -12,6 +12,8 @@ pub(crate) fn compile_point_function(
types: SpirVTypes,
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
.begin_function(
types.float,
@@ -20,7 +22,7 @@ pub(crate) fn compile_point_function(
spirv::FunctionControl::INLINE
| spirv::FunctionControl::PURE
| spirv::FunctionControl::CONST,
types.point_fn_type,
point_fn_type,
)
.unwrap();
let pos_p = b.function_parameter(types.vec4p).unwrap();
@@ -35,7 +37,7 @@ pub(crate) fn compile_point_function(
use SSAOpcode::*;
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(
float: u32,
+2 -10
View File
@@ -467,7 +467,7 @@ impl SSATape {
b
};
let glsl = if with_module {
1
4
} else {
b.ext_inst_import("GLSL.std.450")
};
@@ -480,10 +480,6 @@ impl SSATape {
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 gradient_fn_type = b.type_function(vec4, vec![vec4p]);
let jit_string = b.string("JIT");
const POINT_FN_LOC: u32 = 10000;
const INTERVAL_FN_LOC: u32 = 20000;
@@ -499,10 +495,6 @@ impl SSATape {
vec3,
vec4,
vec4p,
point_fn_type,
interval_fn_type,
gradient_fn_type,
jit_string,
};
// Manually fix the header bounds
@@ -561,7 +553,7 @@ impl SSATape {
if DUMP_SPV_DIS_TO_FILE {
std::fs::write(
format!(
"{}.spv-dis",
"spv-dis/{}.spv-dis",
humantime::format_rfc3339(std::time::SystemTime::now())
),
module.disassemble(),