Skip to content

Commit aa38739

Browse files
committed
perf: the epoch's columns go to the card once
Four things read the same trace columns — the commitment, the sumcheck's factors, the evaluation at the reduction point, and the opening's message — and each uploaded its own copy. Counted on the real block that is 160.9 GiB crossing the bus and 10.21 s of a 44.8 s prove, for forty gigabytes of data. At thirteen to nineteen gigabytes a second they were already running at what pageable host memory gives, so there was nothing to make faster: only three passes to stop making. The columns now go up once, when the epoch's tables are committed, and live as long as the tables that read them. A table's columns are a contiguous run of equal height, which is the layout the factor gather and the batched evaluation already want, so those two read them where they lie; the commitment and the opening scatter theirs by offset, which on the card is a copy at device bandwidth. Every entry point still takes columns that are only here, which is what runs when there is no device or it will not promise the room. Real block, 19 epochs of 2^21: 51.16 -> 45.31 s. Peak VRAM 30223 -> 31055 MiB of 32607 — 832 more, against the 1280 that holding a copy without removing the uploads cost, because what each site used to allocate to upload into is gone. Peak host RSS unchanged at 9.3 GB. The margin this leaves ties the epoch to 2^21: a sweep already measured 32.1 GiB of peak at 2^22 without any of this. Epoch size stops being a free knob and starts depending on a memory model that does not exist yet. Six consecutive runs of the GPU suite green, the cross-check sweep green, the four kernel parity tests green, 274 and 261 on the host.
1 parent c2d73ae commit aa38739

14 files changed

Lines changed: 610 additions & 82 deletions

File tree

‎crypto/math-cuda/src/columns.rs‎

Lines changed: 148 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,148 @@
1+
//! The epoch's trace columns, on the card once.
2+
//!
3+
//! Four things read the same columns — the commitment, the sumcheck's factors,
4+
//! the evaluation at the reduction point, and the opening's message — and each
5+
//! used to upload its own copy. That is the same forty gigabytes crossing the
6+
//! bus four times, and at the speed pageable host memory gives it is a fifth of
7+
//! the prove. They are put here once and read where they lie.
8+
//!
9+
//! A table's columns are a contiguous run of equal height, which is the layout
10+
//! the factor gather and the batched evaluation already want; the commitment and
11+
//! the opening scatter theirs, which on the card is a copy at device bandwidth.
12+
13+
use std::sync::Arc;
14+
15+
use cudarc::driver::{CudaSlice, CudaStream, DevicePtr, DevicePtrMut};
16+
17+
use crate::Result;
18+
use crate::device::{DeviceReservation, alloc_or_trim, backend};
19+
20+
/// Where a set of columns is: still here, or already there.
21+
pub enum Columns<'a> {
22+
Host(&'a [&'a [u64]]),
23+
/// A run of `width` columns of `rows` each, starting at column `first`.
24+
Device {
25+
store: &'a DeviceColumns,
26+
first: usize,
27+
width: usize,
28+
},
29+
}
30+
31+
impl Columns<'_> {
32+
pub fn width(&self) -> usize {
33+
match self {
34+
Self::Host(columns) => columns.len(),
35+
Self::Device { width, .. } => *width,
36+
}
37+
}
38+
39+
pub fn rows(&self) -> usize {
40+
match self {
41+
Self::Host(columns) => columns.first().map_or(0, |c| c.len()),
42+
Self::Device { store, first, .. } => store.spans[*first].1,
43+
}
44+
}
45+
46+
pub fn is_empty(&self) -> bool {
47+
self.width() == 0
48+
}
49+
}
50+
51+
/// The columns themselves, laid end to end in one allocation.
52+
pub struct DeviceColumns {
53+
stream: Arc<CudaStream>,
54+
buffer: CudaSlice<u64>,
55+
/// `(offset, len)` in elements, per column, in the order uploaded.
56+
spans: Vec<(usize, usize)>,
57+
_room: DeviceReservation,
58+
}
59+
60+
impl DeviceColumns {
61+
/// `None` when the card will not promise the room, in which case every
62+
/// caller uploads its own copy as before.
63+
pub fn upload(columns: &[&[u64]]) -> Option<Self> {
64+
if columns.is_empty() {
65+
return None;
66+
}
67+
let total: usize = columns.iter().map(|c| c.len()).sum();
68+
let be = backend().ok()?;
69+
let room = be.reserve(total as u64 * 8)?;
70+
let stream = be.next_stream();
71+
// SAFETY: every element is written by the copies below.
72+
let mut buffer = unsafe { alloc_or_trim::<u64>(&stream, total) }.ok()?;
73+
let mut spans = Vec::with_capacity(columns.len());
74+
let mut at = 0usize;
75+
for column in columns {
76+
let mut slab = buffer.slice_mut(at..at + column.len());
77+
stream.memcpy_htod(*column, &mut slab).ok()?;
78+
spans.push((at, column.len()));
79+
at += column.len();
80+
}
81+
stream.synchronize().ok()?;
82+
Some(Self {
83+
stream,
84+
buffer,
85+
spans,
86+
_room: room,
87+
})
88+
}
89+
90+
pub fn num_columns(&self) -> usize {
91+
self.spans.len()
92+
}
93+
94+
/// Whether `width` columns from `first` are a run of equal height — which
95+
/// is what the kernels that read a table's columns in place need.
96+
pub fn is_run(&self, first: usize, width: usize) -> bool {
97+
if width == 0 || first + width > self.spans.len() {
98+
return false;
99+
}
100+
let (start, rows) = self.spans[first];
101+
(0..width).all(|k| self.spans[first + k] == (start + k * rows, rows))
102+
}
103+
104+
/// The device address of column `k` and its length in elements.
105+
pub fn at(&self, k: usize) -> (u64, usize) {
106+
let (offset, len) = self.spans[k];
107+
let (base, _guard) = self.buffer.device_ptr(&self.stream);
108+
(base + (offset * 8) as u64, len)
109+
}
110+
111+
/// A view over `width` columns from `first`, which must be a run.
112+
pub fn view(&self, first: usize, width: usize) -> cudarc::driver::CudaView<'_, u64> {
113+
let (offset, rows) = self.spans[first];
114+
self.buffer.slice(offset..offset + width * rows)
115+
}
116+
117+
pub fn stream(&self) -> &Arc<CudaStream> {
118+
&self.stream
119+
}
120+
121+
/// Copies column `k` into `dst` at `offset`, on the card.
122+
///
123+
/// Issued on the **destination's** stream, so whatever reads `dst` next is
124+
/// ordered behind it. The store is written once and synchronized at upload
125+
/// and read-only after, so another stream reading it races with nothing.
126+
pub fn copy_into(
127+
&self,
128+
k: usize,
129+
dst: &mut CudaSlice<u64>,
130+
offset: usize,
131+
stream: &Arc<CudaStream>,
132+
) -> Result<()> {
133+
let (src, len) = self.at(k);
134+
let (base, _guard) = dst.device_ptr_mut(stream);
135+
// SAFETY: both ranges are inside allocations this call holds a handle
136+
// to, and the caller checked `offset + len` against `dst`.
137+
unsafe {
138+
cudarc::driver::sys::cuMemcpyDtoDAsync_v2(
139+
base + (offset * 8) as u64,
140+
src,
141+
len * 8,
142+
stream.cu_stream(),
143+
)
144+
.result()?;
145+
}
146+
Ok(())
147+
}
148+
}

‎crypto/math-cuda/src/lib.rs‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@
66
//! pipelines or used by the parity test suite.
77
88
pub mod barycentric;
9+
pub mod columns;
910
pub mod constraint_interp;
1011
pub mod deep;
1112
pub mod device;

‎crypto/math-cuda/src/sumcheck.rs‎

Lines changed: 59 additions & 34 deletions
Original file line numberDiff line numberDiff line change
@@ -553,49 +553,63 @@ pub fn evaluate_mle_base(table: &[u64], point: &[u64]) -> Result<[u64; 3]> {
553553
///
554554
/// `columns` are base-field slices of `2^point.len()/3` values each. Returns
555555
/// one ext3 value per column, in order.
556-
pub fn evaluate_many_base(columns: &[&[u64]], point: &[u64]) -> Result<Vec<[u64; 3]>> {
556+
pub fn evaluate_many_base(
557+
columns: crate::columns::Columns<'_>,
558+
point: &[u64],
559+
) -> Result<Vec<[u64; 3]>> {
557560
assert!(point.len().is_multiple_of(3), "three u64 per coordinate");
558561
let vars = point.len() / 3;
559562
let rows = 1usize << vars;
560563
assert!(vars > 0, "a point with no coordinates is the table itself");
561-
assert!(
562-
columns.iter().all(|column| column.len() == rows),
563-
"every column spans the point"
564-
);
565564
if columns.is_empty() {
566565
return Ok(Vec::new());
567566
}
567+
assert_eq!(columns.rows(), rows, "every column spans the point");
568+
let num_columns = columns.width();
568569

569570
let be = backend()?;
570571
let stream = be.next_stream();
571572
// How many go up at a time: the base copy plus the ext3 half it folds to
572573
// is 20 bytes a row, and this keeps that transient bounded however wide
573574
// the table is.
574575
let per_column = rows as u64 * 20;
575-
let chunk = (CHUNK_BUDGET_BYTES / per_column.max(1)).clamp(1, columns.len() as u64) as usize;
576+
let chunk = (CHUNK_BUDGET_BYTES / per_column.max(1)).clamp(1, num_columns as u64) as usize;
576577

577-
let mut out = Vec::with_capacity(columns.len());
578-
for group in columns.chunks(chunk) {
578+
let mut out = Vec::with_capacity(num_columns);
579+
for start in (0..num_columns).step_by(chunk) {
580+
let group = start..(start + chunk).min(num_columns);
581+
let group_len = group.len();
579582
let half = rows / 2;
580583
// Promised before it is used, like everything else that takes a slab
581584
// of the card: the caller's fallback is to evaluate the columns one at
582585
// a time, which needs almost nothing.
583-
let Some(_room) = crate::device::reserve(group.len() as u64 * per_column) else {
586+
let Some(_room) = crate::device::reserve(group_len as u64 * per_column) else {
584587
return Err(cudarc::driver::DriverError(
585588
cudarc::driver::sys::CUresult::CUDA_ERROR_OUT_OF_MEMORY,
586589
));
587590
};
588-
let mut base = unsafe { alloc_or_trim::<u64>(&stream, group.len() * rows) }?;
589-
for (k, column) in group.iter().enumerate() {
590-
let at = k * rows;
591-
let mut slab = base.slice_mut(at..at + rows);
592-
stream.memcpy_htod(*column, &mut slab)?;
593-
}
591+
// Read where they lie when they are already there; a copy otherwise.
592+
let uploaded;
593+
let base = match &columns {
594+
crate::columns::Columns::Device { store, first, .. } => {
595+
store.view(first + group.start, group_len)
596+
}
597+
crate::columns::Columns::Host(host) => {
598+
let mut up = unsafe { alloc_or_trim::<u64>(&stream, group_len * rows) }?;
599+
for (k, column) in host[group.clone()].iter().enumerate() {
600+
let at = k * rows;
601+
let mut slab = up.slice_mut(at..at + rows);
602+
stream.memcpy_htod(*column, &mut slab)?;
603+
}
604+
uploaded = up;
605+
uploaded.slice(0..group_len * rows)
606+
}
607+
};
594608
let r = crate::device::htod_or_trim(&stream, &point[..3])?;
595609
// SAFETY: the kernel writes every element of the halves it produces.
596-
let mut values = unsafe { alloc_or_trim::<u64>(&stream, group.len() * half * 3) }?;
610+
let mut values = unsafe { alloc_or_trim::<u64>(&stream, group_len * half * 3) }?;
597611
let half_arg = half as u64;
598-
let tables = group.len() as u64;
612+
let tables = group_len as u64;
599613
let total = half_arg * tables;
600614
let grid = total.div_ceil(BLOCK_DIM as u64).clamp(1, MAX_GRID as u64) as u32;
601615
unsafe {
@@ -612,13 +626,11 @@ pub fn evaluate_many_base(columns: &[&[u64]], point: &[u64]) -> Result<Vec<[u64;
612626
shared_mem_bytes: 0,
613627
})?;
614628
}
615-
drop(base);
616-
617629
// From here every table is ext3 and they all fold together: one launch
618630
// per level, over the list of addresses.
619631
let addresses: Vec<u64> = {
620632
let (at, _guard) = values.device_ptr(&stream);
621-
(0..group.len())
633+
(0..group_len)
622634
.map(|k| at + (k * half * 3 * 8) as u64)
623635
.collect()
624636
};
@@ -631,7 +643,7 @@ pub fn evaluate_many_base(columns: &[&[u64]], point: &[u64]) -> Result<Vec<[u64;
631643
break;
632644
}
633645
stream.memcpy_htod(coordinate, &mut r_dev)?;
634-
let width = group.len() as u64;
646+
let width = group_len as u64;
635647
let total = width * fold_half;
636648
let grid = total.div_ceil(BLOCK_DIM as u64).clamp(1, MAX_GRID as u64) as u32;
637649
unsafe {
@@ -656,7 +668,7 @@ pub fn evaluate_many_base(columns: &[&[u64]], point: &[u64]) -> Result<Vec<[u64;
656668
// Gathered there and brought back in one copy. A trace has thousands of
657669
// columns, and reading each one's head on its own is a transfer and a
658670
// stream synchronize apiece for twenty-four bytes.
659-
let heads: Vec<u32> = (0..group.len()).map(|k| (k * half) as u32).collect();
671+
let heads: Vec<u32> = (0..group_len).map(|k| (k * half) as u32).collect();
660672
let packed = crate::fri::gather_ext3_at(&values, &heads, &stream)?;
661673
for head in packed.chunks_exact(3) {
662674
out.push([head[0], head[1], head[2]]);
@@ -782,7 +794,7 @@ impl DeviceFactors {
782794
/// reduced mod `rows`), and the slot it fills; `public` is the extension
783795
/// tables that are not views of a column, each with the slot it goes to.
784796
pub fn from_columns(
785-
columns: &[&[u64]],
797+
columns: crate::columns::Columns<'_>,
786798
plan: &[u64],
787799
public: &[(usize, &[u64])],
788800
rows: usize,
@@ -800,7 +812,7 @@ impl DeviceFactors {
800812
"every slot is filled once"
801813
);
802814
assert!(
803-
columns.iter().all(|column| column.len() == rows),
815+
columns.is_empty() || columns.rows() == rows,
804816
"every column spans the cube"
805817
);
806818

@@ -826,13 +838,26 @@ impl DeviceFactors {
826838
}
827839

828840
if !plan.is_empty() {
829-
// SAFETY: every cell is written by the copies below.
830-
let mut base = unsafe { alloc_or_trim::<u64>(&stream, columns.len() * rows) }?;
831-
for (k, column) in columns.iter().enumerate() {
832-
let at = k * rows;
833-
let mut slab = base.slice_mut(at..at + rows);
834-
stream.memcpy_htod(*column, &mut slab)?;
835-
}
841+
// Read where they lie when they are already there; a copy otherwise.
842+
let uploaded;
843+
let base = match &columns {
844+
crate::columns::Columns::Device {
845+
store,
846+
first,
847+
width,
848+
} => store.view(*first, *width),
849+
crate::columns::Columns::Host(host) => {
850+
// SAFETY: every cell is written by the copies below.
851+
let mut up = unsafe { alloc_or_trim::<u64>(&stream, host.len() * rows) }?;
852+
for (k, column) in host.iter().enumerate() {
853+
let at = k * rows;
854+
let mut slab = up.slice_mut(at..at + rows);
855+
stream.memcpy_htod(*column, &mut slab)?;
856+
}
857+
uploaded = up;
858+
uploaded.slice(0..host.len() * rows)
859+
}
860+
};
836861
let plan_dev = crate::device::htod_or_trim(&stream, plan)?;
837862
let num_plan = (plan.len() / 3) as u64;
838863
let rows_arg = rows as u64;
@@ -853,9 +878,9 @@ impl DeviceFactors {
853878
.arg(&mut buffer)
854879
.launch(cfg)?;
855880
}
856-
// The columns are spent. Freeing them is stream-ordered, so it
857-
// happens behind the kernel that just read them.
858-
drop(base);
881+
// A copy made here is spent and goes at the end of this block,
882+
// which is stream-ordered behind the kernel that just read it; a
883+
// view of the epoch's columns frees nothing, because they stay.
859884
}
860885

861886
let addresses: Vec<u64> = {

‎crypto/math-cuda/src/whir.rs‎

Lines changed: 36 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -193,6 +193,29 @@ pub fn commit_codeword_parts(
193193
)
194194
}
195195

196+
/// The same for parts the card already holds: `(column index, offset)` into the
197+
/// epoch's columns. The scatter is then a copy at device bandwidth rather than
198+
/// the trace crossing the bus again.
199+
pub fn commit_codeword_resident(
200+
store: &crate::columns::DeviceColumns,
201+
parts: &[(usize, usize)],
202+
log_evals: usize,
203+
log_blowup: usize,
204+
log_folding: usize,
205+
transient: bool,
206+
) -> Result<(DeviceCodeword, [u8; 32])> {
207+
commit_from(
208+
Source::Resident {
209+
store,
210+
parts,
211+
log_evals,
212+
},
213+
log_blowup,
214+
log_folding,
215+
transient,
216+
)
217+
}
218+
196219
/// Where a commit's coefficients come from: one slab the host holds, or the
197220
/// columns a stacked polynomial is made of.
198221
enum Source<'a> {
@@ -201,13 +224,18 @@ enum Source<'a> {
201224
parts: &'a [(&'a [u64], usize)],
202225
log_evals: usize,
203226
},
227+
Resident {
228+
store: &'a crate::columns::DeviceColumns,
229+
parts: &'a [(usize, usize)],
230+
log_evals: usize,
231+
},
204232
}
205233

206234
impl Source<'_> {
207235
fn log_evals(&self) -> u64 {
208236
match self {
209237
Self::Whole(evals) => evals.len().trailing_zeros() as u64,
210-
Self::Parts { log_evals, .. } => *log_evals as u64,
238+
Self::Parts { log_evals, .. } | Self::Resident { log_evals, .. } => *log_evals as u64,
211239
}
212240
}
213241

@@ -225,6 +253,13 @@ impl Source<'_> {
225253
}
226254
Ok(())
227255
}
256+
Self::Resident { store, parts, .. } => {
257+
stream.memset_zeros(coeffs)?;
258+
for (column, offset) in *parts {
259+
store.copy_into(*column, coeffs, *offset, stream)?;
260+
}
261+
Ok(())
262+
}
228263
}
229264
}
230265
}

0 commit comments

Comments
 (0)