Data-parallel encrypted batches
Example source: examples/21_distributed_batch_inputs.py
This example starts with one encrypted batch of independent samples. The data owner assigns each rank a consecutive sample interval. Ranks evaluate the same public affine model on their sub-batches, then return their distinct outputs for decryption and concatenation in original sample order.
This is a centralized-input data-parallel workload. Other applications can provision rank-local inputs directly and do not need an input scatter.
Run
python examples/21_distributed_batch_inputs.py --batch-size 7
torchrun --standalone --nproc-per-node=2 \
examples/21_distributed_batch_inputs.py --batch-size 72
3
4
--batch-size is the global number of samples, not a per-rank size. The example requires at least one sample per participating rank. Seven samples on two ranks produce sub-batches of three and four samples. --preset selects the CKKS configuration. dist.init() uses rank-local CUDA devices when available and CPU otherwise; Tensor factories follow the selected PyTorch default device.
1. Encrypt one batch on the data-owner rank
Rank zero constructs a clear Tensor with shape [batch_size, sample_width], creates one encryption key pair, and encrypts the entire batch:
if rank == 0:
secret_key = engine.create_secret_key()
public_key = engine.create_public_key(secret_key)
encrypted_batch = engine.encrypt_message(messages, public_key)2
3
4
The resulting Ciphertext has batch_shape=(batch_size,). Worker ranks need no keys for this plaintext-ciphertext model. Encryption and decryption remain on rank zero.
2. Choose sample intervals and create local views
The application uses a near-even consecutive partition:
batch_ranges = [
(owner * batch_size // world_size,
(owner + 1) * batch_size // world_size)
for owner in range(world_size)
]
if rank == 0:
chunks = [
encrypted_batch.slice_batch(start, stop)
for start, stop in batch_ranges
]
else:
chunks = None2
3
4
5
6
7
8
9
10
11
12
13
slice_batch(start, stop, dim=0) indexes a logical batch axis, excluding the component and RNS axes. It retains the selected axis, including a one-item interval, and accepts negative dim values. Intervals are nonempty half-open ranges with nonnegative bounds. The result shares Tensor storage and preserves depth, scale, prime IDs, and representation state.
The same local slice interface is available on Plaintext and CompressedPlaintext when an application needs corresponding plaintext sub-batches. This example instead uses one shared, unbatched model weight.
The interval list is application policy. Neither slice_batch nor the collective chooses a partition or infers sample identity.
3. Scatter sub-batches
local_input = dist.scatter_ciphertexts(chunks, src=0)
weight = dist.broadcast_plaintext(root_weight, src=0)2
Each entry in chunks is a batched ciphertext; it need not contain the same number of samples as the other entries. The existing typed scatter transmits their shapes and state. It also handles communication packing when the source batch views are non-contiguous. No specialized batch-scatter API is required.
4. Evaluate the same model on every rank
For each sample
The coefficients do not depend on rank or partition. Engine operations process the entire local batch; there is no Python loop over local samples:
local_output = engine.rescale_to_next_depth(
engine.ntt_domain_to_coefficient_domain(
engine.multiply_plaintext(
engine.coefficient_domain_to_ntt_domain(local_input), weight
)
)
)
local_output = engine.add_plaintext(local_output, bias)2
3
4
5
6
7
8
The shared weight is broadcast. Each rank encodes the same bias using the result's depth and actual scale. Multiplication and rescale remain separate operations.
5. Gather and restore sample order
outputs = dist.gather_ciphertexts(local_output, dst=0)The source receives sub-batches in process-group-rank order. Since the chosen intervals are consecutive in that order, rank zero can decrypt each sub-batch and concatenate the clear results along their batch axis:
decoded = torch.cat(
[
engine.decrypt_message(output, secret_key=secret_key, is_real=True)
.cpu()[..., :sample_width]
for output in outputs
],
dim=0,
)2
3
4
5
6
7
8
The reconstructed clear Tensor has shape [batch_size, sample_width] and is checked against the same affine model applied to the original input batch. An arithmetic reduction would add different samples together and change the workload meaning.
Related usage models
- Homogeneous batching compares single-device batched execution with per-sample execution. This example distributes sub-batches across ranks instead.
- Rotation-parallel matrix-vector multiplication distributes contributions to one result and therefore uses arithmetic reduction.
- Communication semantics distinguishes independent samples, additive partials, and RNS shards.
Source
#!/usr/bin/env python3
"""Distribute an encrypted input batch across data-parallel ranks.
Run on one process or one process per GPU:
python examples/21_distributed_batch_inputs.py --batch-size 7
torchrun --standalone --nproc-per-node=2 \
examples/21_distributed_batch_inputs.py --batch-size 7
Rank 0 encrypts one batch, chooses consecutive sample intervals, and scatters
Ciphertext.slice_batch views. Each rank applies the same public affine model
to its entire local batch. Rank 0 gathers and decrypts the sub-batches, then
concatenates their clear results in original sample order. Weights are public;
encryption and decryption keys remain on rank 0.
The partition may be uneven. Neither the model nor sample generation depends
on rank, so changing the partition does not change the intended result.
"""
from __future__ import annotations
import argparse
import torch
from common import parse_preset
import fhelium as fh
import fhelium.distributed as dist
from fhelium.eager import Engine
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--batch-size", type=int, default=8)
parser.add_argument(
"--preset",
type=parse_preset,
default=fh.Preset.slots8192_scale40_depth7_int64,
)
args = parser.parse_args()
dist.init()
try:
rank, world_size = dist.get_rank(), dist.get_world_size()
if args.batch_size < world_size:
raise ValueError(
"batch-size must be at least world size so every rank receives "
"a nonempty sub-batch"
)
torch.set_default_device(dist.local_device())
engine = Engine(args.preset, allow_automatic_key_generation=False)
sample_width = 32
weight_value, bias_value = 1.25, -0.003
batch_ranges = [
(
owner * args.batch_size // world_size,
(owner + 1) * args.batch_size // world_size,
)
for owner in range(world_size)
]
if rank == 0:
secret_key = engine.create_secret_key()
public_key = engine.create_public_key(secret_key)
samples = torch.arange(args.batch_size, dtype=torch.float64)
messages = torch.linspace(
-0.015, 0.015, sample_width, dtype=torch.float64
).unsqueeze(0) + 0.004 * torch.sin(samples * 0.7).unsqueeze(1)
encrypted_batch = engine.encrypt_message(messages, public_key)
chunks = [
encrypted_batch.slice_batch(start, stop)
for start, stop in batch_ranges
]
root_weight = engine.prepare_plaintext_for_multiplication(
engine.encode(
torch.full(
(sample_width,), weight_value, dtype=torch.float64
),
depth=0,
)
)
else:
secret_key = None
messages = None
chunks = None
root_weight = None
# Each list entry is a sub-batch, not a single sample. The collective
# preserves its shape and state; the application chose its sample range.
local_input = dist.scatter_ciphertexts(chunks, src=0)
weight = dist.broadcast_plaintext(root_weight, src=0)
# Engine operations evaluate every item in the local batch. This model
# is independent of both rank number and the chosen batch partition.
local_output = engine.rescale_to_next_depth(
engine.ntt_domain_to_coefficient_domain(
engine.multiply_plaintext(
engine.coefficient_domain_to_ntt_domain(local_input), weight
)
)
)
bias = engine.prepare_plaintext_for_addition(
engine.encode(
torch.full((sample_width,), bias_value, dtype=torch.float64),
depth=local_output.depth,
scale=local_output.scale,
)
)
local_output = engine.add_plaintext(local_output, bias)
outputs = dist.gather_ciphertexts(local_output, dst=0)
start, stop = batch_ranges[rank]
print(
f"rank={rank} samples=[{start},{stop}) "
f"local_batch={local_input.batch_shape[0]} depth={local_output.depth}"
)
if rank == 0:
assert secret_key is not None
assert messages is not None
assert outputs is not None
decoded = torch.cat(
[
engine.decrypt_message(
output, secret_key=secret_key, is_real=True
).cpu()[..., :sample_width]
for output in outputs
],
dim=0,
)
expected = (weight_value * messages + bias_value).cpu()
torch.testing.assert_close(decoded, expected, atol=3e-5, rtol=0)
max_error = float(torch.max(torch.abs(decoded - expected)))
print(
f"data_parallel_batch_ok batch_size={args.batch_size} "
f"world_size={world_size} max_abs_error={max_error:.3e}"
)
finally:
dist.shutdown()
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