TensorPlay
latest (dev)
Copy
View Markdown

Latest development documentation · Updated 2026-10-08

tensorplay.distributed.algorithms.ddp_comm_hooks.default_hooks.fp16_compress_wrapper

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

Cast input tensor to float16, cast result of hook back to input dtype.

This wrapper casts the input gradient tensor of a given DDP communication hook to half-precision floating point format (float16), and casts the resulting tensor of the given hook back to the input data type, such as float32. Therefore, fp16_compress_hook is equivalent to fp16_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, fp16_compress_wrapper(powerSGD_hook))

On this page

Ask DeepWiki