Files
wehub-resource-sync 2aaeece67c
Codestyle Check / Lint (push) Has been cancelled
Codestyle Check / Check bypass (push) Has been cancelled
Pipelines-Test / Pipelines-Test (push) Has been cancelled
chore: import upstream snapshot with attribution
2026-07-13 13:37:14 +08:00

246 lines
10 KiB
Python

# Copyright (c) 2025 PaddlePaddle Authors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import numpy as np
import paddle
import paddle.distributed as dist
from paddle.distributed import fleet
from paddle.distributed.fleet.base import topology
from paddle.distributed.fleet.base.topology import (
CommunicateTopology,
HybridCommunicateGroup,
)
from paddle.distributed.fleet.layers.mpu.random import get_rng_state_tracker
from paddlenlp.transformers.model_utils import unwrap_model
from paddlenlp.utils.log import logger
class ReshardController:
def __init__(
self,
tensor_parallel_degree,
pipeline_parallel_degree=1,
sharding_parallel_degree=1,
sep_parallel_degree=1,
seed=100,
):
self.tensor_parallel_degree = tensor_parallel_degree
self.pipeline_parallel_degree = pipeline_parallel_degree
self.sharding_parallel_degree = sharding_parallel_degree
self.sep_parallel_degree = sep_parallel_degree
self.seed = seed
self.orig_rng_state = paddle.get_rng_state()
self.orig_cuda_rng_state = self._get_rng_state()
self.train_hcg = fleet.get_hybrid_communicate_group()
self.train_tp_group, self.train_dp_group, self.train_sdp_group = (
self.train_hcg.get_model_parallel_group(),
self.train_hcg.get_data_parallel_group(),
self.train_hcg.get_sharding_parallel_group(),
)
self.infer_tp_group, self.infer_dp_group, self.infer_sdp_group = self.init_rollout_env()
self.set_train_env()
self.is_train = True
def init_rollout_env(self):
world_size = dist.get_world_size()
infer_topo = CommunicateTopology(
hybrid_group_names=["data", "pipe", "sharding", "sep", "model"],
dims=[
world_size
// self.tensor_parallel_degree
// self.pipeline_parallel_degree
// self.sharding_parallel_degree
// self.sep_parallel_degree,
self.pipeline_parallel_degree,
self.sharding_parallel_degree,
self.sep_parallel_degree,
self.tensor_parallel_degree,
],
)
infer_hcg = HybridCommunicateGroup(infer_topo)
infer_tp_group = infer_hcg.get_model_parallel_group()
infer_dp_group = infer_hcg.get_data_parallel_group()
infer_sdp_group = infer_hcg.get_sharding_parallel_group()
return (infer_tp_group, infer_dp_group, infer_sdp_group)
def _get_rng_state(self):
"""get_rng_state"""
origin_rng_state = paddle.get_cuda_rng_state()
paddle.seed(self.seed)
rng_state = paddle.get_cuda_rng_state()
paddle.set_cuda_rng_state(origin_rng_state)
return rng_state
def set_rollout_env(self, msg=""):
if "model_parallel_rng" not in get_rng_state_tracker().states_:
local_seed = 2025 + 1 + self.train_hcg.get_model_parallel_rank()
get_rng_state_tracker().add("model_parallel_rng", local_seed)
paddle.set_rng_state(self.orig_cuda_rng_state)
hcg = fleet.get_hybrid_communicate_group()
hcg.get_model_parallel_group = lambda: self.infer_tp_group
hcg.get_model_parallel_world_size = lambda: self.infer_tp_group.nranks
hcg.get_model_parallel_rank = lambda: self.infer_tp_group.rank
hcg.get_sharding_parallel_group = lambda: self.infer_sdp_group
hcg.get_data_parallel_group = lambda: self.infer_dp_group
topology._HYBRID_PARALLEL_GROUP.get_model_parallel_group = lambda: self.infer_tp_group
topology._HYBRID_PARALLEL_GROUP.get_model_parallel_world_size = lambda: self.infer_tp_group.nranks
topology._HYBRID_PARALLEL_GROUP.get_model_parallel_rank = lambda: self.infer_tp_group.rank
self.log(msg, False)
# self.is_train = False
def set_train_env(self, msg=""):
hcg = fleet.get_hybrid_communicate_group()
hcg.get_model_parallel_group = lambda: self.train_tp_group
hcg.get_model_parallel_world_size = lambda: self.train_tp_group.nranks
hcg.get_model_parallel_rank = lambda: self.train_tp_group.rank
hcg.get_sharding_parallel_group = lambda: self.train_sdp_group
hcg.get_data_parallel_group = lambda: self.train_dp_group
topology._HYBRID_PARALLEL_GROUP.get_model_parallel_group = lambda: self.train_tp_group
topology._HYBRID_PARALLEL_GROUP.get_model_parallel_world_size = lambda: self.train_tp_group.nranks
topology._HYBRID_PARALLEL_GROUP.get_model_parallel_rank = lambda: self.train_tp_group.rank
paddle.set_rng_state(self.orig_rng_state)
self.log(msg, True)
# self.is_train = True
def log(self, msg, is_train=False):
msg = f"for {msg}" if len(msg) > 0 else ""
if is_train:
logger.warning(
f"Recover train env done {msg}. [Global TP]: {fleet.get_hybrid_communicate_group().get_model_parallel_world_size()}, [Train TP]: {self.train_tp_group.nranks}, [Infer TP]: {self.infer_tp_group.nranks}"
)
else:
logger.warning(
f"Set rollout env done {msg}. [Global TP]: {fleet.get_hybrid_communicate_group().get_model_parallel_world_size()}, [Train TP]: {self.train_tp_group.nranks}"
)
@paddle.no_grad()
def pp_reshard(tgt_tensor, src_model_state_dict, src_tensor_meta_info, pp_rank, pp_group):
src_tensor_key = src_tensor_meta_info["pipeline_key"]
src_tensor_pp_rank = src_tensor_meta_info["pipeline_src_rank"]
src_tensor_shape = src_tensor_meta_info["shape"]
if src_tensor_pp_rank == pp_rank:
src_tensor = src_model_state_dict.pop(src_tensor_key)
resharded_tensor = src_tensor.clone()
cpu_src_tensor = src_tensor.pin_memory()
cpu_src_tensor._share_buffer_to(src_tensor)
else:
resharded_tensor = paddle.empty(src_tensor_shape)
resharded_tensor = resharded_tensor.astype(tgt_tensor.dtype)
dist.broadcast(resharded_tensor, src=pp_group.ranks[src_tensor_pp_rank], group=pp_group, sync_op=True)
return resharded_tensor
@paddle.no_grad()
def mp_reshard(
src_tensor,
tgt_tensor,
meta_dict,
train_tp_group,
rollout_tp_group,
):
if rollout_tp_group.nranks == train_tp_group.nranks:
return src_tensor
if meta_dict["is_distributed"]:
res = []
if train_tp_group.nranks > 1:
paddle.distributed.all_gather(res, src_tensor, group=train_tp_group, sync_op=True)
else:
res = [src_tensor]
if hasattr(tgt_tensor, "is_distributed") and tgt_tensor.is_distributed:
assert hasattr(tgt_tensor, "split_axis"), f"{tgt_tensor.name} has no split_axis!"
concat_tensor = paddle.concat(res, meta_dict["split_axis"])
del res
all_parts = paddle.split(concat_tensor, rollout_tp_group.nranks, tgt_tensor.split_axis)
del concat_tensor
return all_parts[rollout_tp_group.rank]
else:
return paddle.concat(res, meta_dict["split_axis"])
return src_tensor
def init_reshard_mappings(model, training_args, pp_rank, pp_group):
global_meta_dict = {}
if training_args.pipeline_parallel_degree > 1:
model._layers._set_pipeline_name_mapping()
local_name_mapping_dict = model._layers._single_to_pp_mapping
else:
local_name_mapping_dict = {}
for k in model.state_dict():
local_name_mapping_dict[k] = k.replace("_layers.", "")
local_model_state_dict = unwrap_model(model).state_dict()
local_meta_dict = {}
for k, v in local_name_mapping_dict.items():
if training_args.pipeline_parallel_degree == 1:
k = k.replace("_layers.", "")
pipeline_key = v
pipeline_tensor = local_model_state_dict[pipeline_key]
local_meta_dict[k] = {
"pipeline_key": pipeline_key,
"pipeline_src_rank": pp_rank,
"shape": pipeline_tensor.shape,
}
local_meta_dict[k]["is_distributed"] = False
if hasattr(pipeline_tensor, "is_distributed"):
local_meta_dict[k]["is_distributed"] = pipeline_tensor.is_distributed
local_meta_dict[k]["split_axis"] = None
if hasattr(pipeline_tensor, "split_axis"):
local_meta_dict[k]["split_axis"] = pipeline_tensor.split_axis
if training_args.pipeline_parallel_degree > 1:
gathered_local_meta_dict = []
dist.all_gather_object(gathered_local_meta_dict, local_meta_dict, group=pp_group)
else:
gathered_local_meta_dict = [local_meta_dict]
for meta_dict in gathered_local_meta_dict:
global_meta_dict.update(meta_dict)
return global_meta_dict
@paddle.no_grad()
def reshard_to_rollout(
train_model, rollout_model, global_meta_dict, pp_rank, pp_group, rollout_tp_group, train_tp_group
):
train_model_state_dict = train_model.state_dict()
rollout_model_state_dict = rollout_model.state_dict()
param_numel = [(k, np.prod(v.shape)) for k, v in rollout_model_state_dict.items()]
param_numel.sort(key=lambda x: x[1], reverse=True)
for k, _ in param_numel:
v = rollout_model_state_dict[k]
resharded_tensor = pp_reshard(v, train_model_state_dict, global_meta_dict[k], pp_rank, pp_group)
resharded_tensor = mp_reshard(
resharded_tensor,
v,
global_meta_dict[k],
train_tp_group,
rollout_tp_group,
)
assert resharded_tensor.dtype == v.dtype, f"dtype wrong {k} {resharded_tensor.dtype} {v.dtype}"
assert resharded_tensor.shape == v.shape, f"shape wrong {k} {resharded_tensor.shape} {v.shape}"
resharded_tensor._share_buffer_to(v)
resharded_tensor._clear()
missing_keys = train_model_state_dict.keys()
num_missing_keys = len(missing_keys)
assert num_missing_keys == 0, f"missing {num_missing_keys} keys after reshard policy: {missing_keys}"
logger.warning("[Reshard] Done")