Choose a multi-GPU partition
Choose a partition from the mathematical relationship among rank-local values, then evaluate communication, keys, memory, and load balance. Do not infer a partition from ciphertext shape alone.
Prerequisites
Have a single-rank calculation with known packing, key inventory, depth/scale schedule, and a clear-message reference. The rank-local calculation may use Eager, a linked Program, or a prepared callable; partition its mathematical work before selecting rank-local execution optimizations.
1. Start from a correct single-rank evaluator
Keep a synchronized single-rank implementation with a cleartext oracle. Record its:
- operation and rotation counts;
- depth schedule;
- required keys;
- latency by major phase;
- peak allocated/reserved memory;
- decrypt error.
This establishes the baseline mathematical result and execution cost.
2. Identify independent work
Use this decision tree:
3. Evaluate data parallelism
Choose independent ciphertext data parallelism when each rank owns a separate request/sample.
Plan:
Advantages:
- low communication during evaluation;
- no ciphertext reduction;
- simple scaling and failure localization.
Check whether weights/keys are replicated and whether root encryption or output gather becomes the bottleneck.
4. Evaluate additive-term parallelism
Choose additive-term parallelism when:
Plan:
Advantages:
- expensive rotations/key switches can be partitioned;
- communication may occur mainly in the start and end phases.
Costs:
- input replication;
- per-rank complete active-row layout for local rotations;
- key placement and transient key movement;
- final typed reduction;
- imbalance when term costs differ.
5. Evaluate limb parallelism cautiously
Choose limb parallelism only when a single value/key is too large and the program contains enough row-local work to amortize scatter/gather barriers.
Choose consecutive intervals on the source's stored limb axis and pass them as limb_ranges to scatter_ciphertext_limbs. The interface slices the source and its declared prime IDs; the application determines the partition. Perform only operations with documented partial-layout semantics, and reconstruct every expected active row before rescale, rotation, key switching, relinearization, or decryption.
Repeated complete-row reconstruction points can erase local row-depth gains.
6. Build a cost table
For each candidate, estimate:
| Cost | Data parallel | Additive-term parallel | Limb parallel |
|---|---|---|---|
| Input communication | Scatter | Broadcast | Scatter limbs |
| Output communication | Gather list | Typed ciphertext reduction | Reconstruct limbs |
| Key replication | Often complete evaluator set | Owned steps/common set | Complete-row owner/stage dependent |
| Local values | Independent full values | Full input + partial | Partial rows |
| Synchronization | Start/end | Start/end plus reduce | Every complete-row reconstruction point |
| Best use | Independent requests | Additive packed work | Long row-local phase |
Include topology and transient memory, not only steady retained bytes.
7. Validate world size one and two
Run the same worker under:
- world size one;
- two ranks with a small problem;
- target rank count and problem;
- uneven work distribution;
- an empty-work rank that still participates in collectives.
Compare every final result to the same oracle and report maximum error per rank or at the reconstructed root.
8. Keep graph capture local
If the local evaluator has a fixed schedule, capture one CudaGraphProgram per rank. Leave dynamic input/key provisioning and typed final reduction eager.
9. Report the complete result
Record:
- partition and ownership rule;
- rank count, GPU topology, and launcher;
- local compute, input/key communication, and final gather/reduce separately;
- maximum and aggregate key memory;
- per-rank and global allocator peaks;
- load imbalance;
- correctness;
- startup and process-group setup.
Verify the outcome
Each rank must own a defined input region or contribution, enter the same collective protocol, and produce the single-rank mathematical result after combination. Compare decoded messages under the same error criterion when the distributed summation schedule changes. See examples/21_distributed_batch_inputs.py through examples/24_distributed_collective_ir.py for input, partial-result, row-shard, and IR collective procedures.