Double-buffered execution
Example source: examples/17_runtime_double_buffer.py
ReusableValueBuffer owns fixed storage for an ordinary Tensor/value tree. Example 17 alternates two CUDA buffers while streaming operation-ready plaintext tiles from pinned host memory. It demonstrates data-transfer ordering, not a performance comparison or an automatic memory-admission policy.
python examples/17_runtime_double_buffer.py --device cuda:0
python examples/17_runtime_double_buffer.py --num-tiles 4 --plaintexts-per-tile 4 --message-size 322
The example keeps its established depth-20, logN-16 CKKS configuration and absolute error threshold. The default uses four tiles of four plaintexts, rather than a multi-GiB all-resident benchmark. Two device tiles occupy about 60 MiB; evaluator resources and temporaries are additional allocations.
Prepare the data and storage
Each tile has a different scalar weight sum, at most 0.125. Distinct tile data makes a missed or incorrectly ordered transfer observable. Host plaintexts are independently pinned, and ReusableValueBuffer.like initializes two independent CUDA storage trees from the first tile. The example retains each buffer's ordinary value view and records its Tensor addresses.
Order transfer and computation
At iteration i:
- Enqueue tile i+1 into the other buffer on the transfer stream.
- Its copy waits for the event marking that buffer's previous reader.
- The compute stream waits on the current tile's CopyHandle.
- Evaluate the current tile using its retained value view.
- Record a read-completion event before that buffer can be overwritten.
CopyHandle.wait_on(stream) orders device work without a host synchronization at each iteration. The application retains the host source values while copies are in flight. Streams are synchronized before releasing buffers.
Check the result
The example validates every tile, checks that target addresses did not change, and prints host-weight and fixed-buffer byte counts. It does not include a second all-resident execution mode, allocator flushing, or timing statistics.
Example 18 uses CUDA Graph capture/replay instead of manually scheduling a streaming workload. Example 19 adds managed placement and leases, which are separate from the fixed storage and copy handles demonstrated here.
Source
#!/usr/bin/env python3
"""Stream pinned-host plaintext tiles through two fixed CUDA value buffers.
The transfer stream prepares tile i+1 while the compute stream evaluates tile i.
A CopyHandle makes computation wait for the incoming data; a CUDA event makes
buffer reuse wait for the preceding reader. No CUDA Graph or admission manager
is needed for this application-owned streaming schedule.
"""
from __future__ import annotations
import argparse
from collections.abc import Sequence
import torch
from common import format_bytes, print_table
import fhelium as fh
from fhelium.eager import Engine
from fhelium.runtime import CopyHandle, ReusableValueBuffer
# Preserve the established CKKS configuration and absolute error criterion.
_WORKLOAD_PRESET = fh.Preset.slots32768_scale40_depth34_int64
_WORKLOAD_DEPTH = 20
_VALIDATION_ATOL = 1e-5
def evaluate_weight_tile(
source_ntt: fh.Ciphertext,
weights: Sequence[fh.Plaintext],
*,
engine: Engine,
) -> fh.Ciphertext:
"""Sum plaintext products in NTT form, then invert and rescale once."""
accumulator = engine.multiply_plaintext(source_ntt, weights[0])
for weight in weights[1:]:
engine.add_(accumulator, engine.multiply_plaintext(source_ntt, weight))
return engine.rescale_to_next_depth(
engine.ntt_domain_to_coefficient_domain(accumulator)
)
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--device", default="cuda:0")
parser.add_argument("--num-tiles", type=int, default=4)
parser.add_argument("--plaintexts-per-tile", type=int, default=4)
parser.add_argument("--message-size", type=int, default=32)
args = parser.parse_args()
device = torch.device(args.device)
if device.type != "cuda":
parser.error("Double-buffered asynchronous transfer requires CUDA")
if args.num_tiles < 3 or args.plaintexts_per_tile < 1:
parser.error("Use at least three tiles and one plaintext per tile")
torch.set_default_device(device)
engine = Engine(_WORKLOAD_PRESET, allow_automatic_key_generation=False)
if not 1 <= args.message_size <= engine.num_slots:
parser.error(f"message-size must be between 1 and {engine.num_slots}")
secret_key = engine.create_secret_key()
public_key = engine.create_public_key(secret_key)
message = torch.linspace(
-0.01, 0.01, args.message_size, dtype=torch.float64
)
source = engine.encrypt_message(message, public_key, depth=_WORKLOAD_DEPTH)
source_ntt = engine.coefficient_domain_to_ntt_domain(source)
# Different tile values make missing or incorrectly ordered copies visible.
factors = [0.125 * (i + 1) / args.num_tiles for i in range(args.num_tiles)]
host_tiles = []
for factor in factors:
prototype = engine.prepare_plaintext_for_multiplication(
engine.encode(
factor / args.plaintexts_per_tile, depth=_WORKLOAD_DEPTH
)
).cpu()
host_tiles.append(
[prototype.pin_memory() for _ in range(args.plaintexts_per_tile)]
)
buffers = [
ReusableValueBuffer.like(host_tiles[0], device=device) for _ in range(2)
]
views = [buffer.value for buffer in buffers]
pointers = [
[weight.data.data_ptr() for weight in tile if weight.data is not None]
for tile in views
]
transfer_stream = torch.cuda.Stream(device=device)
compute_stream = torch.cuda.current_stream(device)
transfer_stream.wait_stream(compute_stream)
read_done: list[torch.cuda.Event | None] = [None, None]
current_copy: CopyHandle | None = None
outputs = []
try:
for index in range(args.num_tiles):
current = index % 2
next_copy = None
if index + 1 < args.num_tiles:
next_buffer = (index + 1) % 2
next_copy = buffers[next_buffer].copy_from(
host_tiles[index + 1],
stream=transfer_stream,
non_blocking=True,
wait_for=read_done[next_buffer],
)
if current_copy is not None:
current_copy.wait_on(compute_stream)
outputs.append(
evaluate_weight_tile(source_ntt, views[current], engine=engine)
)
event = torch.cuda.Event()
event.record(compute_stream)
read_done[current] = event
current_copy = next_copy
compute_stream.synchronize()
assert pointers == [
[
weight.data.data_ptr()
for weight in tile
if weight.data is not None
]
for tile in views
]
rows = []
for index, (output, factor) in enumerate(
zip(outputs, factors, strict=True)
):
actual = engine.decrypt_message(output, secret_key, is_real=True)[
: args.message_size
]
expected = message * factor
torch.testing.assert_close(
actual, expected, atol=_VALIDATION_ATOL, rtol=0
)
rows.append([index, factor, float((actual - expected).abs().max())])
print_table(["tile", "weight sum", "maximum error"], rows)
print(
f"Pinned weights: {format_bytes(sum(weight.nbytes for tile in host_tiles for weight in tile))}"
)
print(
f"Two fixed CUDA buffers: {format_bytes(sum(buffer.nbytes for buffer in buffers))}"
)
finally:
transfer_stream.synchronize()
compute_stream.synchronize()
for buffer in buffers:
buffer.close()
if __name__ == "__main__":
main()2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152