flashinfer.comm.PcieIpcReduceScatterWorkspace

class flashinfer.comm.PcieIpcReduceScatterWorkspace(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 SUM reduce-scatter over PCIe IPC.

max_numel is the largest rank-local output shard. Input has shape [world_size * local_rows, hidden] and each rank receives the reduced shard at its rank-major offset with shape [local_rows, hidden]. Input and a caller-provided output must be 16-byte aligned.

BF16, FP16 and FP32 are supported. Each kernel accumulates in FP32 and converts to the requested output dtype. The TP8 variants also convert one four-rank partial before the final accumulation. Results are therefore not promised to be bitwise identical to NCCL.

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.

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, ...])

destroy()

Collectively wait for users, dispose the handle, and free the slab.

launch_config(inp)

Return a deterministic config for inp, or None if unsupported.

rebind_stream()

Allow the next call to bind after the previous stream is ordered.

reduce_scatter(inp, *[, out, config])

Reduce the rank-major input and return this rank's output shard.

supports(inp)

Whether this workspace can reduce-scatter inp.

tune(hiddens, *[, cache, warmup, repeat])

Measure configs for hiddens over this workspace's tune batches.

tuned_launch_config(inp)

Return a measured config when cached, otherwise the seed config.

Attributes

device

dtype

element_size

group

handle

max_blocks

max_numel

ordered_4plus4

ordered_4plus4_reason

placement_fingerprint

rank

world_size