[Hardware] AMD - Fix checkpoint saving or conversion (#301)

This commit is contained in:
gramesh-amd
2025-12-05 12:30:25 +08:00
committed by GitHub
parent 6c34a6efbe
commit 4f6b95f02e
5 changed files with 44 additions and 74 deletions
+2 -1
View File
@@ -70,8 +70,9 @@ Use [mbridge](https://github.com/ISEEKYAN/mbridge.git) or [Megatron-LM-amd_versi
cd miles/
source scripts/models/qwen3-4B.sh
MEGATRON_LM_PATH=$(pip list | grep megatron-core | awk '{print $NF}')
PYTHONPATH=${MEGATRON_LM_PATH} python tools/convert_hf_to_torch_dist_amd.py \
PYTHONPATH=${MEGATRON_LM_PATH} python tools/convert_hf_to_torch_dist.py \
${MODEL_ARGS[@]} \
--no-gradient-accumulation-fusion \
--hf-checkpoint model/Qwen3-4B \
--save model/Qwen3-4B_torch_dist
```
+7
View File
@@ -714,6 +714,13 @@ def initialize_model_and_optimizer(
tuple[list[DDP], MegatronOptimizer, OptimizerParamScheduler, int]:
DDP-wrapped model chunks, optimizer, scheduler, and iteration index.
"""
if torch.version.hip:
import megatron.core.dist_checkpointing.strategies.filesystem_async as filesystem_async_module
from miles.utils.rocm_checkpoint_writer import ROCmFileSystemWriterAsync
filesystem_async_module.FileSystemWriterAsync = ROCmFileSystemWriterAsync
print("[ROCm] Applied FileSystemWriterAsync patch for HIP compatibility")
model, optimizer, opt_param_scheduler = setup_model_and_optimizer(args, role)
model[0].role = role
clear_memory()
+29
View File
@@ -0,0 +1,29 @@
import torch
from megatron.core.dist_checkpointing.strategies.filesystem_async import (
FileSystemWriterAsync,
)
class ROCmFileSystemWriterAsync(FileSystemWriterAsync):
"""
FileSystemWriterAsync wrapper for ROCm compatibility.
On ROCm/HIP, using non_blocking=True causes tensors to be stored in pinned memory,
which triggers segmentation faults when forking subprocesses afterward.
"""
@staticmethod
def preload_tensors(*args, **kwargs):
# Change argument non_blocking to False on HIP platform
# The tensors will be stored in pinned memory if non_blocking=True
# Currently on the ROCm platform, forking a subprocess afterward
# with pinned_memory=True will trigger segmentation fault
if torch.version.hip:
print("HIP/ROCm detected: setting non_blocking=False in preload_tensors")
if "non_blocking" in kwargs:
kwargs["non_blocking"] = False
elif len(args) > 1 and isinstance(args[-1], bool):
# non_blocking is typically the last argument
args = args[:-1] + (False,)
return FileSystemWriterAsync.preload_tensors(*args, **kwargs)
+6
View File
@@ -72,6 +72,12 @@ def get_args():
def main():
if torch.version.hip:
import megatron.core.dist_checkpointing.strategies.filesystem_async as filesystem_async_module
from miles.utils.rocm_checkpoint_writer import ROCmFileSystemWriterAsync
filesystem_async_module.FileSystemWriterAsync = ROCmFileSystemWriterAsync
print("[ROCm] Applied FileSystemWriterAsync patch for HIP compatibility")
configure_logger()
# Initialize distributed environment
-73
View File
@@ -1,73 +0,0 @@
import os
import shutil
import torch
import torch.distributed as dist
from megatron.core import parallel_state as mpu
from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed
from megatron.training.arguments import parse_args
from megatron.training.checkpointing import get_checkpoint_name, get_checkpoint_tracker_filename, save_checkpoint
from megatron.training.global_vars import set_args
from mbridge import AutoBridge
def init_distributed():
"""Initialize distributed environment"""
os.environ["RANK"] = "0"
os.environ["WORLD_SIZE"] = "1"
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "12357"
backend = "gloo"
print("Backend:", backend)
torch.distributed.init_process_group(backend)
mpu.initialize_model_parallel(
tensor_model_parallel_size=1,
virtual_pipeline_model_parallel_size=None,
context_parallel_size=1,
expert_model_parallel_size=1,
)
model_parallel_cuda_manual_seed(0)
def add_convertion_args(parser):
"""Add conversion arguments to the parser"""
parser.add_argument("--hf-checkpoint", type=str, required=True, help="HuggingFace model path")
return parser
def main():
# Parse command line arguments
args = parse_args(add_convertion_args)
args.use_dist_ckpt = args.ckpt_format != "torch"
set_args(args)
# Initialize distributed environment
init_distributed()
# Load model
hf_model_path = args.hf_checkpoint
bridge = AutoBridge.from_pretrained(hf_model_path)
model = bridge.get_model(use_cpu_initialization=True)
bridge.load_weights(model, hf_model_path)
print(f"Model loaded: {hf_model_path}")
model[0].config.use_cpu_initialization = True
model[0] = model[0].cpu()
save_checkpoint(1, model, None, None, 0)
if dist.get_rank() == 0:
# change to release ckpt
tracker_filename = get_checkpoint_tracker_filename(args.save)
with open(tracker_filename, "w") as f:
f.write("release")
source_dir = get_checkpoint_name(args.save, 1, False, return_base_dir=True)
target_dir = get_checkpoint_name(args.save, -1, True, return_base_dir=True)
shutil.move(source_dir, target_dir)
dist.barrier()
dist.destroy_process_group()
if __name__ == "__main__":
main()