# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the BSD-style license found in the
# LICENSE file in the root directory of this source tree.

from __future__ import annotations

import logging
import math
import os
from collections.abc import Iterable
from datetime import timedelta
from typing import TYPE_CHECKING

import torch
import torch.distributed._functional_collectives as funcol
import torch.distributed.config as dist_config
import torch.distributed.distributed_c10d as c10d
import torch.distributed.tensor._random
import torch.distributed.tensor.parallel
from torch import distributed as dist
from torch.distributed.device_mesh import DeviceMesh
from torch.distributed.tensor import DTensor

from torchtitan.config import CommConfig, DebugConfig
from torchtitan.distributed.parallelism_context import DistributedTopology
from torchtitan.tools.utils import device_module, device_type, get_local_device

logger = logging.getLogger(__name__)


if TYPE_CHECKING:
    from torchtitan.distributed.parallelism_context import ParallelismContext


def _dist_reduce(
    x: torch.Tensor,
    reduceOp: str,
    mesh: DeviceMesh | None,
    extra_pg: dist.ProcessGroup | None,
) -> float:
    """Perform distributed reduction on a tensor.

    Args:
        x (torch.Tensor): Input tensor.
        reduceOp (str): Reduce operation to perform.
        mesh (DeviceMesh | None): Device mesh to use for reduction.
            If None, no reduction is performed but simply convert the tensor to a float.
        extra_pg (dist.ProcessGroup, optional): Extra process group to use for reduction.
            Defaults to None. If provided, this all_reduce will be called for the extra
            process group, and then the result will be all_reduced for the mesh.
    """
    return float(_dist_reduce_tensor(x, reduceOp, mesh, extra_pg).item())


def _dist_reduce_tensor(
    x: torch.Tensor,
    reduceOp: str,
    mesh: DeviceMesh | None,
    extra_pg: dist.ProcessGroup | None,
) -> torch.Tensor:
    """Perform a distributed reduction without moving the result to the CPU."""
    needs_wait = False
    if extra_pg is not None:
        x = funcol.all_reduce(x, reduceOp=reduceOp, group=extra_pg)
        needs_wait = True
    if mesh is not None:
        x = funcol.all_reduce(x, reduceOp=reduceOp, group=mesh)
        needs_wait = True
    return funcol.wait_tensor(x) if needs_wait else x


# TODO: rename this to maybe_dist_max
def dist_max(
    x: torch.Tensor,
    mesh: DeviceMesh | None = None,
    extra_pg: dist.ProcessGroup | None = None,
) -> float:
    return _dist_reduce(
        x, reduceOp=c10d.ReduceOp.MAX.name, mesh=mesh, extra_pg=extra_pg
    )


def dist_sum(
    x: torch.Tensor,
    mesh: DeviceMesh | None = None,
    extra_pg: dist.ProcessGroup | None = None,
) -> float:
    return _dist_reduce(
        x, reduceOp=c10d.ReduceOp.SUM.name, mesh=mesh, extra_pg=extra_pg
    )


def dist_sum_tensor(
    x: torch.Tensor,
    mesh: DeviceMesh | None = None,
    extra_pg: dist.ProcessGroup | None = None,
) -> torch.Tensor:
    """Sum a tensor across process groups and keep the result on its device."""
    return _dist_reduce_tensor(
        x, reduceOp=c10d.ReduceOp.SUM.name, mesh=mesh, extra_pg=extra_pg
    )


def set_determinism(
    parallelism_context: ParallelismContext,
    device: torch.device,
    debug_config: DebugConfig,
    distinct_seed_mesh_axes: list[str],
) -> None:
    """
    Set the same distributed RNG seed for all axes in the world mesh, but use
    different seeds across axes named by ``distinct_seed_mesh_axes``. For
    example, pipeline stages should use different seeds while ranks within an
    SPMD group use the same seed.

    This uses PyTorch's DTensor RNG tracker because it provides mesh-aware RNG
    offsets for sharded parameter initialization.

    Set Determinism flags for increased reproducibility with loss of performance.

    Args:
        parallelism_context: Parallelism context for distributed training.
        device: Device to use
        debug_config: Debug config to use
        distinct_seed_mesh_axes: Mesh axis names that receive distinct seeds.
    """
    if debug_config.deterministic:
        logger.info("Deterministic algorithm enabled (expect perf degradation).")
        torch.use_deterministic_algorithms(
            True, warn_only=debug_config.deterministic_warn_only
        )
        torch.backends.cudnn.deterministic = True
        torch.backends.cudnn.benchmark = False
        # use_deterministic_algorithms(True) enables fill_uninitialized_memory,
        # which makes torch.empty() run a fill kernel. This kernel races with
        # DeepEP comm streams, causing errors.
        # This also prevents HF modeling from initializing ROPE (inv_freq) buffers to NaN.
        # pyrefly: ignore [missing-attribute]
        torch.utils.deterministic.fill_uninitialized_memory = False
        # env var for deterministic CuBLAS
        # https://pytorch.org/docs/stable/generated/torch.use_deterministic_algorithms.html
        os.environ["CUBLAS_WORKSPACE_CONFIG"] = ":4096:8"

        from torch.nn.attention.flex_attention import flex_attention

        from torchtitan.models.common.attention import FlexInnerAttention

        if torch.version.hip is not None:
            # Compiled ROCm flex attention is not deterministic.
            # Falling back to eager (non-compiled) flex_attention for determinism on ROCm.
            logger.info(
                "Using eager (non-compiled) flex_attention for determinism on ROCm."
            )
            FlexInnerAttention._compiled_flex_attn = flex_attention
        else:
            # Ensure flex_attention is compiled without max-autotune. This is needed to ensure
            # reproducibility, since the autotune results may not be deterministic. We disable
            # autotune in-place on FlexInnerAttention.inductor_configs (rather than recompiling with
            # no options) so the regional-inductor scoop configs are preserved.
            FlexInnerAttention.inductor_configs["max_autotune"] = False
            FlexInnerAttention.inductor_configs["coordinate_descent_tuning"] = False
            # pyrefly: ignore [no-matching-overload]
            FlexInnerAttention._compiled_flex_attn = torch.compile(
                flex_attention, options=FlexInnerAttention.inductor_configs
            )

    if debug_config.detect_anomaly:
        logger.warning(
            "Anomaly detection enabled. This incurs significant overhead "
            "and is for debugging only."
        )
        # check_nan=False disables the NaN/Inf gradient check that internally calls
        # aten._is_any_true, which has no DTensor sharding strategy and would crash.
        # Stack trace recording (the useful part) is still enabled.
        torch.autograd.set_detect_anomaly(True, check_nan=False)

    seed = debug_config.seed
    if parallelism_context.world_size == 1:
        if seed is not None:
            torch.manual_seed(seed)
            os.environ["PYTHONHASHSEED"] = str(seed % 2**32)
            logger.debug(f"Single-process job using seed: {seed}")
        return

    # to ensure we can control which ranks have same or different seeds, all ranks agree on a starting seed.
    # if user provides one, we use this. Otherwise rank 0 rolls the dice and everyone else uses that.
    if seed is None:
        # Extract the seed for torch's main generator on rank 0 and standardizes on using that to build
        # seeds for unique SPMD groups
        seed_tensor = torch.get_rng_state()[:8].to(device)
        torch.distributed.broadcast(seed_tensor, src=0)
        seed = seed_tensor.to("cpu").view(torch.uint64).item()
    assert isinstance(seed, int)

    # Set distinct seeds across the requested mesh axes.
    # For PP + SPMD cases, we want to separate the world into the SPMD mesh and the PP mesh,
    # and choose a unique seed for each rank on the PP mesh.
    # We support multiple distinct dimensions by adding each distinct dimension's local rank to the seed.
    distinct_seed_meshes = [
        parallelism_context.get_optional_mesh(axis) for axis in distinct_seed_mesh_axes
    ]
    distinct_seed_meshes = [mesh for mesh in distinct_seed_meshes if mesh is not None]
    assert all(mesh is not None for mesh in distinct_seed_meshes)

    if distinct_seed_meshes:
        # Each axis contributes: local_rank * (product of all previous axis sizes).
        # This guarantees uniqueness like multi-dimensional array indexing
        seed_offset = 0
        cumulative_size = 1

        for distinct_mesh in distinct_seed_meshes:
            local_rank = distinct_mesh.get_local_rank()
            # Add this axis's contribution.
            seed_offset += local_rank * cumulative_size
            # Update cumulative size for the next axis.
            cumulative_size *= distinct_mesh.size()

        seed += seed_offset
        seed %= 2**64

        logger.debug(
            f"Distinct axes {distinct_seed_mesh_axes}, Global rank {c10d.get_rank()} using seed: {seed}"
        )

    else:
        logger.debug(f"Global Rank {c10d.get_rank()} using seed: {seed}")

    # The native RNGs and python RNG may not be important, except for the 1-D PP case, but we seed them for consistency.
    torch.manual_seed(seed)
    # PYTHONHASHSEED can be a decimal number in the range [0, 2**32 - 1]
    os.environ["PYTHONHASHSEED"] = str(seed % 2**32)

    # As long as we are not in the 1-D (PP-only) case, we will have a seed to use for
    # all ranks of the SPMD mesh. If PP is also used, this seed is unique per PP rank.
    # TODO: remove the need of passing in a mesh once
    # torch.distributed.tensor._random.manual_seed doesn't require a mesh input.
    if parallelism_context.world_size > parallelism_context.pp:
        # We just need to pass the world_mesh as the device_id is the only information
        # this API uses.
        torch.distributed.tensor._random.manual_seed(
            seed, parallelism_context.world_mesh
        )


def init_fake_mode(
    world_size: int,
    *,
    rank: int = 0,
) -> None:
    """Initialize fake backend

    Args:
        world_size: The number of GPUs to simulate
        rank: Global rank to simulate
    """
    torch.distributed.init_process_group(
        "fake",
        rank=rank,
        world_size=world_size,
    )


def _env_int(name: str, *, default: int | None = None) -> int:
    """Read an integer environment variable with an actionable error."""
    value = os.environ.get(name)
    if value is None:
        if default is not None:
            return default
        raise ValueError(f"{name} environment variable must be set")
    try:
        return int(value)
    except ValueError as error:
        raise ValueError(
            f"{name} environment variable must be a valid integer, got: {value}"
        ) from error


def _fake_logical_rank(logical_world_size: int, pp_degree: int) -> int:
    """Resolve a pure-fake logical rank with SPMD coordinate zero."""
    if logical_world_size % pp_degree != 0:
        raise ValueError(
            f"Logical world size {logical_world_size} must be divisible by PP "
            f"degree {pp_degree}"
        )
    pp_rank = _env_int("FAKE_PP_RANK", default=0 if pp_degree == 1 else None)
    spmd_world_size = logical_world_size // pp_degree
    if not 0 <= pp_rank < pp_degree:
        raise ValueError(f"FAKE_PP_RANK must be in [0, {pp_degree}), got {pp_rank}")
    return pp_rank * spmd_world_size


def _init_real_pp_fake_spmd(
    logical_world_size: int,
    pp_degree: int,
    timeout: timedelta,
) -> DistributedTopology:
    """Initialize real PP communication inside a fake logical SPMD world."""
    physical_world_size = _env_int("WORLD_SIZE")
    physical_rank = _env_int("RANK")
    if physical_world_size != pp_degree:
        raise ValueError(
            "real-PP/fake-SPMD mode requires one physical process per PP rank: "
            f"WORLD_SIZE={physical_world_size}, PP={pp_degree}"
        )
    if not 0 <= physical_rank < physical_world_size:
        raise ValueError(
            f"RANK must be in [0, {physical_world_size}), got {physical_rank}"
        )
    if logical_world_size % pp_degree != 0:
        raise ValueError(
            f"Logical world size {logical_world_size} must be divisible by PP "
            f"degree {pp_degree}"
        )

    if "FAKE_PP_RANK" in os.environ:
        raise ValueError(
            "FAKE_PP_RANK is invalid with the real_pp_fake_spmd backend; "
            "physical RANK selects the PP coordinate"
        )
    spmd_world_size = logical_world_size // pp_degree
    logical_rank = physical_rank * spmd_world_size
    logical_pp_ranks = [pp_rank * spmd_world_size for pp_rank in range(pp_degree)]
    init_fake_mode(logical_world_size, rank=logical_rank)

    rendezvous = dist.rendezvous(
        "env://",
        rank=physical_rank,
        world_size=physical_world_size,
        timeout=timeout,
    )
    store, rendezvous_rank, rendezvous_world_size = next(rendezvous)
    if (rendezvous_rank, rendezvous_world_size) != (
        physical_rank,
        physical_world_size,
    ):
        raise RuntimeError("Physical PP rendezvous returned inconsistent topology")

    group_name = c10d.GroupName("torchtitan_real_pp")
    pp_group, _ = c10d._new_process_group_helper(
        group_size=physical_world_size,
        group_rank=physical_rank,
        global_ranks_in_group=logical_pp_ranks,
        backend="nccl",
        store=store,
        group_name=group_name,
        timeout=timeout,
        pg_tag=group_name,
        device_id=get_local_device(),
        group_desc="TorchTitan real pipeline group",
    )
    if not isinstance(pp_group, dist.ProcessGroup):
        raise RuntimeError("Failed to construct the real PP process group")
    c10d._world.pg_group_ranks[pp_group] = {
        logical_rank: group_rank
        for group_rank, logical_rank in enumerate(logical_pp_ranks)
    }
    return DistributedTopology(
        world_size=logical_world_size,
        real_pp_group_for_fake_spmd=pp_group,
    )


def init_distributed(
    comm_config: CommConfig,
    enable_cpu_backend: bool = False,
    base_folder: str = "",
    ranks: list[int] | None = None,
    *,
    pipeline_parallel_degree: int = 1,
) -> DistributedTopology:
    """Initialize communication and return the logical distributed topology."""
    # Skip initialization if already initialized
    if torch.distributed.is_initialized():
        logger.warning(
            "torch.distributed is already initialized. Skipping init_distributed. "
            "The provided comm_config and other settings will not take effect."
        )
        return DistributedTopology(torch.distributed.get_world_size())

    # Directed physical PP edges need independent, preinitialized communicator
    # FIFOs so eager execution and CUDA graph replay use deterministic ordering.
    # Older PyTorch versions do not expose this option.
    if hasattr(dist_config, "pipeline_per_edge_p2p"):
        setattr(  # noqa: B010
            dist_config, "pipeline_per_edge_p2p", pipeline_parallel_degree > 1
        )
    elif pipeline_parallel_degree > 1:
        raise RuntimeError(
            "Pipeline parallelism requires a PyTorch version that provides "
            "torch.distributed.distributed_c10d._config.pipeline_per_edge_p2p."
        )

    # disable autograd multithreading, to enable TLS DeviceMesh stack for spmd_types backend.
    # this is needed for AC functionality; multi-threaded autograd means BWD threads performing recompute,
    # cannot access PGs, e.g. current_spmd_mesh().get_group("tp") to perform the collectives they need.
    torch.autograd.set_multithreading_enabled(False)

    if comm_config.backend in {"fake", "real_pp_fake_spmd"}:
        logical_world_size = _env_int("NGPU")
        if comm_config.backend == "real_pp_fake_spmd":
            return _init_real_pp_fake_spmd(
                logical_world_size,
                pipeline_parallel_degree,
                timedelta(seconds=comm_config.init_timeout_seconds),
            )
        rank = _fake_logical_rank(logical_world_size, pipeline_parallel_degree)
        if not 0 <= rank < logical_world_size:
            raise ValueError(
                f"Fake rank must be in [0, {logical_world_size}), got {rank}"
            )
        init_fake_mode(logical_world_size, rank=rank)
        return DistributedTopology(logical_world_size)

    def _warn_overwrite_env(env, val):
        if env in os.environ:
            logger.warning(
                f"ENV[{env}] = {os.environ[env]} will be overridden to {val} based on job config"
            )
        os.environ[env] = val

    def _get_distributed_backend(enable_cpu_backend):
        backend = "nccl"
        if device_type in torch.distributed.Backend.default_device_backend_map:
            backend = torch.distributed.Backend.default_device_backend_map.get(
                device_type
            )
        if enable_cpu_backend:
            backend = f"{device_type}:{backend},cpu:gloo"
        return backend

    TRACE_BUFFER_SIZE = "TORCH_FR_BUFFER_SIZE"
    TRACE_FILE = "TORCH_FR_DUMP_TEMP_FILE"
    DUMP_ON_TIMEOUT = "TORCH_NCCL_DUMP_ON_TIMEOUT"
    ASYNC_ERROR_HANDLING = "TORCH_NCCL_ASYNC_ERROR_HANDLING"
    SKIP_CLEANUP = "3"

    # FlightRecorder is incompatible with =1 mode where watchdog aborts work, must use =3 (skipcleanup)
    # to get flight recorder dumps. See https://github.com/pytorch/pytorch/issues/121055
    # This could be done only when flight recorder is enabled, but its nice to be consistent to avoid subtle
    # behavior differences
    _warn_overwrite_env(ASYNC_ERROR_HANDLING, SKIP_CLEANUP)

    # enable torch nccl flight recorder in the mode that would dump files if timeout is detected
    _warn_overwrite_env(TRACE_BUFFER_SIZE, str(comm_config.trace_buf_size))
    if comm_config.trace_buf_size > 0:
        # dump on timeout by default if trace buffer is enabled
        _warn_overwrite_env(DUMP_ON_TIMEOUT, "1")
        dump_dir = os.path.join(base_folder, comm_config.save_traces_folder)
        prefix = comm_config.save_traces_file_prefix
        os.makedirs(dump_dir, exist_ok=True)
        _warn_overwrite_env(TRACE_FILE, f"{dump_dir}/{prefix}")

    torch.distributed.init_process_group(
        backend=_get_distributed_backend(enable_cpu_backend),
        timeout=timedelta(seconds=comm_config.init_timeout_seconds),
        _ranks=ranks if ranks is not None else [],
    )

    return DistributedTopology(torch.distributed.get_world_size())


def set_pg_timeouts(
    timeout: timedelta,
    parallelism_context: ParallelismContext,
):
    """
    Sets the timeout for all PGs in the provided mesh, and the default (world) group.

    Note: synchronizes via a barrier, before changing the timeouts. This is important, because
    otherwise you may face a race where the slow rank has not reached the timeout reduction point
    yet due to slow operations permitted under the old timeout value, but other faster ranks may
    start issuing collectives under the new shorter timeout and then immediately timeout.
    """
    logger.info(
        f"Synchronizing and adjusting timeout for all ProcessGroups to {timeout}"
    )
    # Ensure that all the ranks have reached the point of setting the new timeout-
    # otherwise, some ranks may issue collectives with the new/shorter timeout and
    # those may time out, before other ranks have finished with initialization done
    # under the old/slow timeout.
    torch.distributed.barrier(device_ids=[device_module.current_device()])
    device_module.synchronize()

    # None represents the 'default' PG, not part of the mesh
    groups: list[torch.distributed.ProcessGroup | None] = [
        mesh.get_group()
        for mesh in parallelism_context.get_all_one_dimensional_meshes().values()
    ] + [None]
    for group in groups:
        torch.distributed.set_timeout(timeout, group)


@torch.no_grad()
def clip_grad_norm_(
    parameters: torch.Tensor | Iterable[torch.Tensor],
    max_norm: float,
    norm_type: float = 2.0,
    error_if_nonfinite: bool = False,
    foreach: bool | None = None,
    pp_mesh: DeviceMesh | None = None,
    ep_enabled: bool = False,
) -> torch.Tensor:
    """
    Clip the gradient norm of an iterable of parameters.

    Gradient norm clipping requires computing the gradient norm over the entire model.
    `torch.nn.utils.clip_grad_norm_` only computes gradient norm along DP/FSDP/TP dimensions.
    We need to manually reduce the gradient norm across PP stages.
    See https://github.com/pytorch/torchtitan/issues/596 for details.

    Args:
        parameters: an iterable of Tensors or a single Tensor that will have gradients normalized
        max_norm (float): max norm of the gradients
        norm_type (float): type of the used p-norm. Can be ``'inf'`` for
            infinity norm.
        error_if_nonfinite (bool): if True, an error is thrown if the total
            norm of the gradients from :attr:`parameters` is ``nan``,
            ``inf``, or ``-inf``. Default: False (will switch to True in the future)
        foreach (bool): use the faster foreach-based implementation.
            If ``None``, use the foreach implementation for CUDA and CPU native tensors and silently
            fall back to the slow implementation for other device types.
            Default: ``None``
        pp_mesh: Pipeline Parallel device mesh. If not None, will reduce gradient norm across PP stages.
        ep_dense_params_mesh_ndim: Mesh ndim of the dense params when EP is used. If EP is not used,
            set it to ``None``.

    Returns:
        Total norm of the parameter gradients (viewed as a single vector).

    """
    if ep_enabled:
        return _clip_grad_norm_with_ep(
            parameters,
            max_norm,
            norm_type,
            error_if_nonfinite,
            foreach,
            pp_mesh,
        )

    if isinstance(parameters, torch.Tensor):
        parameters = [parameters]
    else:
        # prevent generators from being exhausted
        parameters = list(parameters)
    grads = [p.grad for p in parameters if p.grad is not None]
    total_norm = torch.nn.utils.get_total_norm(
        grads, norm_type, error_if_nonfinite, foreach
    )

    # If total_norm is a DTensor, the placements must be `torch.distributed._tensor.ops.math_ops._NormPartial`.
    # We can simply reduce the DTensor to get the total norm in this tensor's process group
    # and then convert it to a local tensor.
    # NOTE: It has two purposes:
    #       1. to make sure the total norm is computed correctly when PP is used (see below)
    #       2. to return a reduced total_norm tensor whose .item() would return the correct value
    if isinstance(total_norm, DTensor):
        # Will reach here if any non-PP parallelism is used.
        # If only using PP, total_norm will be a local tensor.
        total_norm = total_norm.full_tensor()

    if pp_mesh is not None:
        if math.isinf(norm_type):
            dist.all_reduce(total_norm, op=dist.ReduceOp.MAX, group=pp_mesh.get_group())
        else:
            total_norm **= norm_type
            dist.all_reduce(total_norm, op=dist.ReduceOp.SUM, group=pp_mesh.get_group())
            total_norm **= 1.0 / norm_type

    torch.nn.utils.clip_grads_with_norm_(parameters, max_norm, total_norm, foreach)
    return total_norm


@torch.no_grad()
def _clip_grad_norm_with_ep(
    parameters: torch.Tensor | Iterable[torch.Tensor],
    max_norm: float,
    norm_type: float,
    error_if_nonfinite: bool,
    foreach: bool | None,
    pp_mesh: DeviceMesh | None,
) -> torch.Tensor:
    ep_params = []
    non_ep_params = []
    ep_grads = []
    non_ep_grads = []

    for p in parameters:
        if p.grad is None:
            continue
        assert isinstance(p, DTensor) and isinstance(p.grad, DTensor)
        mesh_dim_names = p.device_mesh.mesh_dim_names
        assert mesh_dim_names is not None
        if "ep" in mesh_dim_names:
            ep_params.append(p)
            ep_grads.append(p.grad)
        else:
            non_ep_params.append(p)
            non_ep_grads.append(p.grad)

    # Either list can be empty depending on the parallelization strategy:
    # - In torchtitan with separate dense/sparse meshes, both lists are typically non-empty
    # - In autoparallel, all params may live on a single sparse mesh with "ep" dimension,
    #   so non_ep_grads would be empty
    # - In PP + EP setups, certain PP ranks may only own EP or non-EP layers
    ep_grads_total_norm = torch.nn.utils.get_total_norm(
        ep_grads, norm_type, error_if_nonfinite, foreach
    )
    # get_total_norm returns tensor(0.) for empty list, which is a non-DTensor
    if isinstance(ep_grads_total_norm, DTensor):
        ep_grads_total_norm = ep_grads_total_norm.full_tensor()

    non_ep_grads_total_norm = torch.nn.utils.get_total_norm(
        non_ep_grads, norm_type, error_if_nonfinite, foreach
    )
    # get_total_norm returns tensor(0.) for empty list, which is a non-DTensor
    if isinstance(non_ep_grads_total_norm, DTensor):
        non_ep_grads_total_norm = non_ep_grads_total_norm.full_tensor()

    if math.isinf(norm_type):
        total_norm = torch.maximum(ep_grads_total_norm, non_ep_grads_total_norm)
    else:
        total_norm = (
            ep_grads_total_norm**norm_type + non_ep_grads_total_norm**norm_type
        )
        total_norm **= 1.0 / norm_type

    if pp_mesh is not None:
        if math.isinf(norm_type):
            dist.all_reduce(total_norm, op=dist.ReduceOp.MAX, group=pp_mesh.get_group())
        else:
            total_norm **= norm_type
            dist.all_reduce(total_norm, op=dist.ReduceOp.SUM, group=pp_mesh.get_group())
            total_norm **= 1.0 / norm_type

    torch.nn.utils.clip_grads_with_norm_(ep_params, max_norm, total_norm, foreach)
    torch.nn.utils.clip_grads_with_norm_(non_ep_params, max_norm, total_norm, foreach)

    return total_norm
