#![cfg(feature = "gpu")]
#[test]
fn single_stream_read_ceiling() {
let Some((device, queue)) = pollster::block_on(async {
let inst = wgpu::Instance::new(wgpu::InstanceDescriptor {
backends: wgpu::Backends::VULKAN | wgpu::Backends::METAL,
flags: wgpu::InstanceFlags::default(),
memory_budget_thresholds: Default::default(),
backend_options: wgpu::BackendOptions::default(),
display: None,
});
let adapter = inst
.request_adapter(&wgpu::RequestAdapterOptions {
power_preference: wgpu::PowerPreference::HighPerformance,
..Default::default()
})
.await
.ok()?;
let limits = adapter.limits();
adapter
.request_device(&wgpu::DeviceDescriptor {
required_limits: limits,
..Default::default()
})
.await
.ok()
}) else {
eprintln!("no adapter — skipping");
return;
};
const WGSL: &str = r#"
@group(0) @binding(0) var<storage, read> src: array<vec4<f32>>;
@group(0) @binding(1) var<storage, read_write> dst: array<f32>;
struct P { vecs_per_wg: u32, _a: u32, _b: u32, _c: u32 };
@group(0) @binding(2) var<uniform> p: P;
// 256 lanes stride a contiguous slice: lane i reads vec i, i+256, ...
// so every 16-load wavefront touches one 16 KB run of DRAM.
@compute @workgroup_size(256)
fn stream_sum(@builtin(workgroup_id) wid: vec3<u32>,
@builtin(local_invocation_index) lid: u32) {
let base = wid.x * p.vecs_per_wg;
var acc = vec4<f32>(0.0);
var i = lid;
loop {
if (i >= p.vecs_per_wg) { break; }
acc = acc + src[base + i];
i = i + 256u;
}
if (lid == 0u) { dst[wid.x] = acc.x + acc.y + acc.z + acc.w; }
}
"#;
let module = device.create_shader_module(wgpu::ShaderModuleDescriptor {
label: Some("bw"),
source: wgpu::ShaderSource::Wgsl(WGSL.into()),
});
let pipe = device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
label: Some("stream_sum"),
layout: None,
module: &module,
entry_point: Some("stream_sum"),
compilation_options: Default::default(),
cache: None,
});
let bytes: u64 = 1_750_000_000 / 16 * 16;
let nvec = (bytes / 16) as u32;
let src = device.create_buffer(&wgpu::BufferDescriptor {
label: Some("src"),
size: bytes,
usage: wgpu::BufferUsages::STORAGE,
mapped_at_creation: false,
});
let wgs: u32 = 2048;
let vecs_per_wg = nvec.div_ceil(wgs);
let dst = device.create_buffer(&wgpu::BufferDescriptor {
label: Some("dst"),
size: (wgs * 4) as u64,
usage: wgpu::BufferUsages::STORAGE,
mapped_at_creation: false,
});
let ubuf = device.create_buffer(&wgpu::BufferDescriptor {
label: Some("p"),
size: 16,
usage: wgpu::BufferUsages::UNIFORM | wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
});
queue.write_buffer(&ubuf, 0, bytemuck::cast_slice(&[vecs_per_wg, 0u32, 0, 0]));
let bind = device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipe.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry { binding: 0, resource: src.as_entire_binding() },
wgpu::BindGroupEntry { binding: 1, resource: dst.as_entire_binding() },
wgpu::BindGroupEntry { binding: 2, resource: ubuf.as_entire_binding() },
],
});
let run = || {
let mut enc = device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&pipe);
pass.set_bind_group(0, &bind, &[]);
pass.dispatch_workgroups(wgs, 1, 1);
}
queue.submit([enc.finish()]);
let _ = device.poll(wgpu::PollType::wait_indefinitely());
};
run(); let reps = 5;
let t = std::time::Instant::now();
for _ in 0..reps {
run();
}
let per = t.elapsed().as_secs_f64() / reps as f64;
eprintln!(
"single-stream read: {:.1} GB in {:.2} ms = {:.0} GB/s",
bytes as f64 / 1e9,
per * 1e3,
bytes as f64 / per / 1e9
);
const WGSL16: &str = r#"
@group(0) @binding(0) var<storage, read> src: array<vec4<f32>>;
@group(0) @binding(1) var<storage, read_write> dst: array<f32>;
struct P { vecs_per_row: u32, rows_per_wg: u32, _b: u32, _c: u32 };
@group(0) @binding(2) var<uniform> p: P;
@compute @workgroup_size(256)
fn stream16(@builtin(workgroup_id) wid: vec3<u32>,
@builtin(local_invocation_index) lid: u32) {
// 16 lanes a row, like the matvec: lane's row = lid/16, its
// stride walks the row 16 vec4 at a time.
let row = wid.x * p.rows_per_wg + (lid >> 4u);
let base = row * p.vecs_per_row;
var acc = vec4<f32>(0.0);
var i = lid & 15u;
loop {
if (i >= p.vecs_per_row) { break; }
acc = acc + src[base + i];
i = i + 16u;
}
if (lid == 0u) { dst[wid.x] = acc.x + acc.y + acc.z + acc.w; }
}
"#;
let m16 = device.create_shader_module(wgpu::ShaderModuleDescriptor {
label: Some("bw16"),
source: wgpu::ShaderSource::Wgsl(WGSL16.into()),
});
let p16 = device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
label: Some("stream16"),
layout: None,
module: &m16,
entry_point: Some("stream16"),
compilation_options: Default::default(),
cache: None,
});
let vecs_per_row: u32 = 256; let rows_total = nvec / vecs_per_row;
let rows_per_wg: u32 = 16;
let wgs16 = rows_total / rows_per_wg;
queue.write_buffer(&ubuf, 0, bytemuck::cast_slice(&[vecs_per_row, rows_per_wg, 0u32, 0]));
let bind16 = device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &p16.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry { binding: 0, resource: src.as_entire_binding() },
wgpu::BindGroupEntry { binding: 1, resource: dst.as_entire_binding() },
wgpu::BindGroupEntry { binding: 2, resource: ubuf.as_entire_binding() },
],
});
let run16 = || {
let mut enc = device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&p16);
pass.set_bind_group(0, &bind16, &[]);
pass.dispatch_workgroups(wgs16, 1, 1);
}
queue.submit([enc.finish()]);
let _ = device.poll(wgpu::PollType::wait_indefinitely());
};
run16();
let t = std::time::Instant::now();
for _ in 0..reps {
run16();
}
let per16 = t.elapsed().as_secs_f64() / reps as f64;
eprintln!(
"16-row interleaved read: {:.1} GB in {:.2} ms = {:.0} GB/s",
bytes as f64 / 1e9,
per16 * 1e3,
bytes as f64 / per16 / 1e9
);
let slices: u32 = 320;
let vecs_per_slice = nvec / slices;
queue.write_buffer(&ubuf, 0, bytemuck::cast_slice(&[vecs_per_slice, 0u32, 0, 0]));
let wg_per_slice = 64u32; let run320 = || {
let mut enc = device.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&pipe);
pass.set_bind_group(0, &bind, &[]);
for _ in 0..slices {
pass.dispatch_workgroups(wg_per_slice, 1, 1);
}
}
queue.submit([enc.finish()]);
let _ = device.poll(wgpu::PollType::wait_indefinitely());
};
run320();
let t = std::time::Instant::now();
for _ in 0..reps {
run320();
}
let per320 = t.elapsed().as_secs_f64() / reps as f64;
eprintln!(
"320 serialized dispatches (structure only): {:.2} ms = {:.2} us per dispatch",
per320 * 1e3,
per320 * 1e6 / slices as f64
);
}