# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 import argparse import json import os import pickle import re import shutil import time import safetensors.torch import torch import torch.distributed.checkpoint as dist_cp from transformers import AutoConfig from typing_extensions import override from slime.backends.megatron_utils.megatron_to_hf import convert_to_hf, remove_padding class UnpicklerWrapper(pickle.Unpickler): @override def find_class(self, mod_name, name): class DummyClass: def __init__(self, *args, **kwargs): pass if mod_name.startswith("megatron") or mod_name.startswith("glm"): return DummyClass return super().find_class(mod_name, name) pickle.Unpickler = UnpicklerWrapper class WrappedStorageReader(dist_cp.FileSystemReader): @override def read_metadata(self): path = self.fs.concat_path(self.path, ".metadata") with self.fs.create_stream(path, "rb") as metadata_file: metadata = UnpicklerWrapper(metadata_file).load() if getattr(metadata, "storage_meta", None) is None: metadata.storage_meta = dist_cp.StorageMeta() metadata.storage_meta.load_id = self.load_id if metadata.planner_data is None: metadata.planner_data = {} return metadata class EmptyStateDictLoadPlanner(dist_cp.default_planner.DefaultLoadPlanner): @override def set_up_planner( self, state_dict: dist_cp.metadata.STATE_DICT_TYPE, metadata: dist_cp.metadata.Metadata | None = None, is_coordinator: bool = False, ) -> None: for k, v in metadata.state_dict_metadata.items(): if "optimizer" in k or "_state" in k: continue print(f"find {k} in torch_dist ckpt") if isinstance(v, dist_cp.metadata.TensorStorageMetadata): v = torch.empty(v.size, dtype=v.properties.dtype) # type: ignore[assignment] state_dict[k] = v super().set_up_planner(state_dict, metadata, is_coordinator) def get_expert_param(args, name, param): if ".experts." not in name: yield name, param return num_experts = args.num_experts match = re.search(r"mlp.experts\.(.+)\.weight(\d+)", name) if not match: assert param.shape[0] == num_experts for expert_id in range(num_experts): expert_name = name.replace(".experts.experts.", ".experts.") + str(expert_id) expert_param = param[expert_id] yield expert_name, expert_param else: yield name, param def get_layer_param(args, name, param): if ".layers." not in name: yield name, param return num_layers = args.num_layers match = re.search(r"\.layers\.(\d+)\.", name) if not match: assert param.shape[0] == num_layers for layer_id in range(num_layers): layer_name = name.replace(".layers.", f".layers.{layer_id}.") layer_param = param[layer_id] yield from get_expert_param(args, layer_name, layer_param) else: yield from get_expert_param(args, name, param) def get_named_params(args, state_dict): for name, param in state_dict.items(): name = f"module.module.{name}" yield from get_layer_param(args, name, param) def save_tensors(args, model_name, state_dict, output_dir, chunk_size, vocab_size=None): # for slime update_weight compatible args.sglang_enable_ep_moe = False print(f"start saving to {output_dir}") os.makedirs(output_dir, exist_ok=True) # 2GB current_size = 0 total_size = 0 modeltensors = [{}] for name, param in get_named_params(args, state_dict): if vocab_size: param = remove_padding(name, param, vocab_size) converted_named_tensors = convert_to_hf(args, model_name, name, param) for converted_name, converted_param in converted_named_tensors: tensor_size = converted_param.numel() * converted_param.element_size() if tensor_size + current_size > chunk_size: modeltensors.append({}) current_size = 0 modeltensors[-1][converted_name] = converted_param current_size += tensor_size total_size += tensor_size metadata = {"metadata": {"total_size": total_size}, "weight_map": {}} num_files = len(modeltensors) for i, tensors in enumerate(modeltensors): filename = f"model-{i:05d}-of-{num_files:05d}.safetensors" for key in tensors.keys(): metadata["weight_map"][key] = filename index_filepath = os.path.join(output_dir, "model.safetensors.index.json") json.dump(metadata, open(index_filepath, "w"), indent=2) print(f"{index_filepath} saved.") for i, tensors in enumerate(modeltensors): filename = f"model-{i:05d}-of-{num_files:05d}.safetensors" t = time.time() filepath = os.path.join(output_dir, filename) safetensors.torch.save_file(tensors, filepath) print(f"{filename} saved in {time.time() - t:.2f} sec.") def copy_assets(origin_hf_dir, output_dir): for filename in os.listdir(origin_hf_dir): if filename == "model.safetensors.index.json" or filename.endswith(".safetensors"): continue origin_filename = os.path.join(origin_hf_dir, filename) if not os.path.isfile(origin_filename): print(f"Skip {filename}, not a file.") continue src, dst = origin_filename, os.path.join(output_dir, filename) print(f"copy from {src} to {dst}") shutil.copy(src, dst) if __name__ == "__main__": parser = argparse.ArgumentParser() parser.add_argument("--model-name", type=str, default=None) parser.add_argument("--input-dir", type=str, required=True) parser.add_argument("--output-dir", type=str, required=True) parser.add_argument( "--origin-hf-dir", type=str, default=None, help="use the origin hf dir to copy files like tokenizer, config.json, etc.", ) parser.add_argument( "-f", "--force", action="store_true", help="Force overwrite the output directory if it exists." ) parser.add_argument( "--chunk-size", type=int, default=5 * 1024**3, help="Chunk size for saving tensors, default is 2GB.", ) parser.add_argument( "--vocab-size", type=int, default=None, help="Vocab size for removing padding, if applicable. If not provided, no padding will be removed.", ) args = parser.parse_args() if os.path.exists(args.output_dir) and not args.force: raise ValueError(f"Output directory {args.output_dir} already exists. Use --force to overwrite it.") if args.model_name is None and args.origin_hf_dir is None: raise ValueError( "Either --model-name or --origin-hf-dir must be provided, so that we can know the name of the params." ) if args.model_name is None: hf_config = AutoConfig.from_pretrained(args.origin_hf_dir, trust_remote_code=True) args.model_name = type(hf_config).__name__.lower() state_dict = {} print(f"loading model from {args.input_dir}") t = time.time() megatron_args = torch.load(os.path.join(args.input_dir, "common.pt"), weights_only=False)["args"] dist_cp.state_dict_loader._load_state_dict( state_dict, storage_reader=WrappedStorageReader(args.input_dir), planner=EmptyStateDictLoadPlanner(), no_dist=True, ) print(f"model loaded in {time.time()-t:.2f} sec.") save_tensors(megatron_args, args.model_name, state_dict, args.output_dir, args.chunk_size, args.vocab_size) if args.origin_hf_dir: copy_assets(args.origin_hf_dir, args.output_dir)