flashinfer.comm.PcieIpcAllGatherWorkspace¶
- class flashinfer.comm.PcieIpcAllGatherWorkspace(group: ProcessGroup, max_numel: int, dtype: dtype = torch.bfloat16, max_blocks: int = 64, tune_batches: Sequence[int] = (1, 2, 4, 8, 16, 32, 64, 128, 256, 512), tune_cache: str | None = None)¶
Workspace for out-of-place PCIe IPC all-gather.
max_numelis the largest rank-local input shard, not the gathered output size. Input has shape[local_rows, hidden]and output is rank-major with shape[world_size * local_rows, hidden]. Input and a caller-provided output must be 16-byte aligned.Construction, calls, and destruction are collective. Every rank must issue the same call sequence and launch configuration. One workspace is bound to one ordered CUDA stream and owns a slab independent from every other PCIe IPC workspace.
A captured CUDA graph keeps using this workspace’s handle and IPC slab. Keep the workspace alive until every replay has finished, and never replay a graph concurrently with another call that uses the same workspace.
BF16, FP16 and FP32 are supported. All variants move opaque 16-byte packs, so the input bit pattern is preserved exactly. The default launch uses a measured exact-shape cache when available and otherwise falls back to a conservative seed.
- __init__(group: ProcessGroup, max_numel: int, dtype: dtype = torch.bfloat16, max_blocks: int = 64, tune_batches: Sequence[int] = (1, 2, 4, 8, 16, 32, 64, 128, 256, 512), tune_cache: str | None = None) None¶
Methods
__init__(group, max_numel[, dtype, ...])all_gather(inp, *[, out, config])Gather one rank-local shard into a rank-major output tensor.
destroy()Collectively wait for users, dispose the handle, and free the slab.
launch_config(inp)Return a deterministic config for
inp, orNoneif unsupported.rebind_stream()Allow the next call to bind after the previous stream is ordered.
supports(inp)Whether this workspace can all-gather
inp.tune(hiddens, *[, cache, warmup, repeat])Measure configs for
hiddensover this workspace's tune batches.tuned_launch_config(inp)Return a measured config when cached, otherwise the seed config.
Attributes
devicedtypeelement_sizegrouphandlemax_blocksmax_numelordered_4plus4ordered_4plus4_reasonplacement_fingerprintrankworld_size