latest (dev)
Copy
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 asfloat32. Therefore,fp16_compress_hookis equivalent tofp16_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))
Help improve this page
Found an error, an unclear step, or a missing example?
Was this page helpful?

