TensorPlay
API symbolsdistributed
Copy
View MarkdownDownload .md

tensorplay.distributed.algorithms.model_averaging.utils.get_params_to_average

tensorplay.distributed.algorithms.model_averaging.utils.get_params_to_average(params: Iterable[Parameter] | Iterable[dict[str, Parameter]])[source]

Return a list of parameters that need to average.

This filters out the parameters that do not contain any gradients. :param params: The parameters of a model or parameter groups of an optimizer.

Ask DeepWiki