use ferrotherm::wgsl::{sweep_shader, GpuModel};
use wgpu::util::DeviceExt;
pub struct Gpu {
device: wgpu::Device,
queue: wgpu::Queue,
info: wgpu::AdapterInfo,
}
impl Gpu {
pub fn new() -> Option<Gpu> {
let instance = wgpu::Instance::new(wgpu::InstanceDescriptor::new_without_display_handle());
let adapter = pollster::block_on(instance.request_adapter(&wgpu::RequestAdapterOptions {
power_preference: wgpu::PowerPreference::HighPerformance,
force_fallback_adapter: false,
compatible_surface: None,
apply_limit_buckets: false,
}))
.ok()?;
let info = adapter.get_info();
let (device, queue) = pollster::block_on(adapter.request_device(&wgpu::DeviceDescriptor {
label: Some("ferrotherm"),
required_features: wgpu::Features::empty(),
required_limits: wgpu::Limits::default(),
memory_hints: wgpu::MemoryHints::Performance,
experimental_features: wgpu::ExperimentalFeatures::disabled(),
trace: wgpu::Trace::Off,
}))
.ok()?;
Some(Gpu { device, queue, info })
}
pub fn adapter(&self) -> &wgpu::AdapterInfo {
&self.info
}
pub fn is_hardware(&self) -> bool {
!matches!(self.info.device_type, wgpu::DeviceType::Cpu | wgpu::DeviceType::Other)
}
pub fn sweep(
&self,
m: &GpuModel,
spins: &mut [i8],
beta: f64,
sweeps: u32,
) -> Result<(), String> {
if spins.len() != m.n as usize {
return Err(format!(
"this model has {} nodes and that state has {}",
m.n,
spins.len()
));
}
if !beta.is_finite() || beta < 0.0 {
return Err(format!("beta must be finite and non-negative, not {beta}"));
}
if m.classes.is_empty() {
return Err("a model with no colour classes has nothing to dispatch".into());
}
let dev = &self.device;
let storage = wgpu::BufferUsages::STORAGE;
let rw = storage | wgpu::BufferUsages::COPY_DST | wgpu::BufferUsages::COPY_SRC;
let pad_u32 = |v: &[u32]| if v.is_empty() { vec![0u32] } else { v.to_vec() };
let pad_f32 = |v: &[f32]| if v.is_empty() { vec![0f32] } else { v.to_vec() };
let mk_u32 = |label: &str, data: &[u32], usage: wgpu::BufferUsages| {
dev.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some(label),
contents: bytes_u32(&pad_u32(data)),
usage,
})
};
let mk_f32 = |label: &str, data: &[f32], usage: wgpu::BufferUsages| {
dev.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some(label),
contents: bytes_f32(&pad_f32(data)),
usage,
})
};
let b_nbr = mk_u32("nbr", &m.nbr, storage);
let b_w = mk_f32("w", &m.w, storage);
let b_h = mk_f32("h", &m.h, storage);
let state: Vec<i32> = spins.iter().map(|&s| s as i32).collect();
let b_spin = dev.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("spin"),
contents: bytes_i32(&state),
usage: rw,
});
let b_dbg = mk_f32("dbg", &vec![0f32; m.n as usize], rw);
let classes: Vec<(u32, wgpu::Buffer)> = m
.classes
.iter()
.map(|c| (c.len() as u32, mk_u32("cls", c, storage)))
.collect();
let readback = dev.create_buffer(&wgpu::BufferDescriptor {
label: Some("readback"),
size: (state.len() * 4) as u64,
usage: wgpu::BufferUsages::COPY_DST | wgpu::BufferUsages::MAP_READ,
mapped_at_creation: false,
});
let module = dev.create_shader_module(wgpu::ShaderModuleDescriptor {
label: Some("sweep"),
source: wgpu::ShaderSource::Wgsl(sweep_shader().into()),
});
let sto = |ro: bool| wgpu::BindingType::Buffer {
ty: wgpu::BufferBindingType::Storage { read_only: ro },
has_dynamic_offset: false,
min_binding_size: None,
};
let entry = |binding: u32, ty: wgpu::BindingType| wgpu::BindGroupLayoutEntry {
binding,
visibility: wgpu::ShaderStages::COMPUTE,
ty,
count: None,
};
let layout = dev.create_bind_group_layout(&wgpu::BindGroupLayoutDescriptor {
label: Some("sweep"),
entries: &[
entry(0, wgpu::BindingType::Buffer {
ty: wgpu::BufferBindingType::Uniform,
has_dynamic_offset: true,
min_binding_size: wgpu::BufferSize::new(PARAMS_BYTES),
}),
entry(1, sto(true)),
entry(2, sto(true)),
entry(3, sto(true)),
entry(4, sto(true)),
entry(5, sto(false)),
entry(6, sto(false)),
],
});
let pipeline_layout = dev.create_pipeline_layout(&wgpu::PipelineLayoutDescriptor {
label: Some("sweep"),
bind_group_layouts: &[Some(&layout)],
immediate_size: 0,
});
let pipeline = dev.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
label: Some("sweep"),
layout: Some(&pipeline_layout),
module: &module,
entry_point: Some("sweep"),
compilation_options: Default::default(),
cache: None,
});
let stride = align_up(PARAMS_BYTES, dev.limits().min_uniform_buffer_offset_alignment as u64);
let live: Vec<usize> = (0..classes.len()).filter(|&i| classes[i].0 > 0).collect();
if live.is_empty() {
return Err("every colour class is empty; there is nothing to sample".into());
}
let steps = sweeps as usize * live.len();
let mut params = vec![0u8; steps * stride as usize];
for s in 0..sweeps as usize {
for (li, &ci) in live.iter().enumerate() {
let step = (s * live.len() + li + 1) as u32;
let at = (s * live.len() + li) * stride as usize;
let p = &mut params[at..at + PARAMS_BYTES as usize];
p[0..4].copy_from_slice(&m.n.to_le_bytes());
p[4..8].copy_from_slice(&m.k.to_le_bytes());
p[8..12].copy_from_slice(&classes[ci].0.to_le_bytes());
p[12..16].copy_from_slice(&step.to_le_bytes());
p[16..20].copy_from_slice(&(beta as f32).to_le_bytes());
}
}
let b_params = dev.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("params"),
contents: ¶ms,
usage: wgpu::BufferUsages::UNIFORM,
});
let binds: Vec<wgpu::BindGroup> = live
.iter()
.map(|&ci| {
dev.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("sweep"),
layout: &layout,
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: wgpu::BindingResource::Buffer(wgpu::BufferBinding {
buffer: &b_params,
offset: 0,
size: wgpu::BufferSize::new(PARAMS_BYTES),
}),
},
wgpu::BindGroupEntry { binding: 1, resource: b_nbr.as_entire_binding() },
wgpu::BindGroupEntry { binding: 2, resource: b_w.as_entire_binding() },
wgpu::BindGroupEntry { binding: 3, resource: b_h.as_entire_binding() },
wgpu::BindGroupEntry { binding: 4, resource: classes[ci].1.as_entire_binding() },
wgpu::BindGroupEntry { binding: 5, resource: b_spin.as_entire_binding() },
wgpu::BindGroupEntry { binding: 6, resource: b_dbg.as_entire_binding() },
],
})
})
.collect();
let mut enc = dev.create_command_encoder(&Default::default());
{
let mut pass = enc.begin_compute_pass(&Default::default());
pass.set_pipeline(&pipeline);
for s in 0..sweeps as usize {
for (li, &ci) in live.iter().enumerate() {
let off = ((s * live.len() + li) * stride as usize) as u32;
pass.set_bind_group(0, &binds[li], &[off]);
pass.dispatch_workgroups(classes[ci].0.div_ceil(WORKGROUP), 1, 1);
}
}
}
enc.copy_buffer_to_buffer(&b_spin, 0, &readback, 0, (state.len() * 4) as u64);
self.queue.submit(Some(enc.finish()));
let slice = readback.slice(..);
let (tx, rx) = std::sync::mpsc::channel();
slice.map_async(wgpu::MapMode::Read, move |r| {
let _ = tx.send(r);
});
self.device.poll(wgpu::PollType::wait_indefinitely()).map_err(|e| format!("device poll failed: {e:?}"))?;
rx.recv()
.map_err(|_| "the readback never completed".to_string())?
.map_err(|e| format!("the readback failed: {e:?}"))?;
{
let data = slice.get_mapped_range().map_err(|e| format!("mapping failed: {e:?}"))?;
for (i, chunk) in data.chunks_exact(4).enumerate().take(spins.len()) {
let v = i32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]);
if v != 1 && v != -1 {
return Err(format!("the GPU returned {v} at spin {i}; states are +1/-1"));
}
spins[i] = v as i8;
}
}
readback.unmap();
Ok(())
}
}
const WORKGROUP: u32 = 64;
const PARAMS_BYTES: u64 = 32;
fn align_up(v: u64, to: u64) -> u64 {
v.div_ceil(to) * to
}
fn bytes_u32(v: &[u32]) -> &[u8] {
unsafe { std::slice::from_raw_parts(v.as_ptr() as *const u8, std::mem::size_of_val(v)) }
}
fn bytes_i32(v: &[i32]) -> &[u8] {
unsafe { std::slice::from_raw_parts(v.as_ptr() as *const u8, std::mem::size_of_val(v)) }
}
fn bytes_f32(v: &[f32]) -> &[u8] {
unsafe { std::slice::from_raw_parts(v.as_ptr() as *const u8, std::mem::size_of_val(v)) }
}
#[cfg(test)]
mod tests {
use super::*;
use ferrotherm::wgsl::GpuModel;
use ferrotherm::gibbs::Sampler;
use ferrotherm::ising::lattice2d;
macro_rules! gpu_or_skip {
() => {
match Gpu::new() {
Some(g) => g,
None => {
eprintln!("no GPU adapter on this machine; skipping");
return;
}
}
};
}
#[test]
fn the_workgroup_size_matches_the_shader() {
let src = ferrotherm::wgsl::sweep_shader();
assert!(
src.contains(&format!("@workgroup_size({WORKGROUP})")),
"this crate dispatches in groups of {WORKGROUP}; the shader says otherwise"
);
}
#[test]
fn a_ferromagnet_orders_at_low_temperature_and_melts_at_high() {
let gpu = gpu_or_skip!();
let g = lattice2d(16, 1.0);
let m = GpuModel::from_graph(&g);
let mag = |beta: f64| {
let mut s = vec![1i8; 256];
gpu.sweep(&m, &mut s, beta, 400).unwrap();
(s.iter().map(|&x| x as f64).sum::<f64>() / 256.0).abs()
};
let cold = mag(1.0);
let hot = mag(0.05);
assert!(cold > 0.8, "a ferromagnet at beta=1 should be ordered, got |m| = {cold:.3}");
assert!(hot < 0.4, "and disordered at beta=0.05, got |m| = {hot:.3}");
}
#[test]
fn the_gpu_reproduces_the_exact_mean_energy() {
let gpu = gpu_or_skip!();
let g = lattice2d(4, 1.0);
let n = 16.0;
let solver = ferrotherm::exact::Elimination { max_width: 20 };
let ln_z = |beta: f64| solver.log_partition(&g, beta).unwrap().log_z.expect("log_partition returns log_z");
let beta = 0.7; let h = 1e-3;
let exact_per_site = -(ln_z(beta + h) - ln_z(beta - h)) / (2.0 * h) / n;
let runs = 24;
let mut total = 0.0;
for r in 0..runs {
let m = GpuModel::from_graph(&g);
let mut s: Vec<i8> = (0..16).map(|i| if (i + r) % 2 == 0 { 1 } else { -1 }).collect();
gpu.sweep(&m, &mut s, beta, 400).unwrap();
total += g.energy(&s);
}
let got = total / runs as f64 / n;
assert!(
(got - exact_per_site).abs() < 0.12,
"GPU {got:.4} vs exact {exact_per_site:.4} per site at beta {beta} -- the shader is \
sampling a different distribution from the one the model defines"
);
}
#[test]
fn the_gpu_and_the_cpu_agree_away_from_criticality() {
let gpu = gpu_or_skip!();
let g = lattice2d(12, 1.0);
let n = 144.0;
let beta = 0.7;
let start: Vec<i8> = (0..144).map(|i| if i % 2 == 0 { 1 } else { -1 }).collect();
let m = GpuModel::from_graph(&g);
let mut s = start.clone();
gpu.sweep(&m, &mut s, beta, 800).unwrap();
let e_gpu = g.energy(&s) / n;
let mut sim = Sampler::new(&g, beta, 7);
sim.s = start;
sim.sweeps(800, None);
let e_cpu = g.energy(&sim.s) / n;
assert!(
(e_gpu - e_cpu).abs() < 0.12,
"GPU {e_gpu:.4} vs CPU {e_cpu:.4} per site -- two implementations of one update rule"
);
}
#[test]
fn a_state_that_is_not_plus_or_minus_one_is_refused_rather_than_coerced() {
let gpu = gpu_or_skip!();
let g = lattice2d(4, 1.0);
let m = GpuModel::from_graph(&g);
let mut wrong = vec![1i8; 9];
let e = gpu.sweep(&m, &mut wrong, 0.5, 1).unwrap_err();
assert!(e.contains("16 nodes") && e.contains('9'), "must name both counts: {e}");
}
#[test]
fn a_bad_temperature_is_refused_by_name() {
let gpu = gpu_or_skip!();
let g = lattice2d(4, 1.0);
let m = GpuModel::from_graph(&g);
let mut s = vec![1i8; 16];
for bad in [f64::NAN, f64::INFINITY, -1.0] {
let e = gpu.sweep(&m, &mut s, bad, 1).unwrap_err();
assert!(e.contains("beta"), "{bad} should be refused by name, got: {e}");
}
}
}