TensorPlay
latest (dev)
Copy
View Markdown

Latest development documentation · Updated 2026-10-08

tensorplay.distributed.algorithms.ddp_comm_hooks.default_hooks.bf16_compress_wrapper

tensorplay.distributed.algorithms.ddp_comm_hooks.default_hooks.bf16_compress_wrapper(hook: Callable[[Any], Any])[source]

Warning: This API is experimental, and it requires NCCL version later than 2.9.6.

This wrapper casts the input gradient tensor of a given DDP communication hook to half-precision Brain floating point format (bfloat16) and casts the resulting tensor of the given hook back to the input data type.

Therefore, bf16_compress_hook is equivalent to bf16_compress_wrapper(allreduce_hook).

Example::
>>> # xdoctest: +SKIP
>>> state = PowerSGDState(process_group=process_group, matrix_approximation_rank=1, start_powerSGD_iter=10)
>>> ddp_model.register_comm_hook(state, bf16_compress_wrapper(powerSGD_hook))

On this page

Ask DeepWiki