mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
This commit is contained in:
@@ -718,9 +718,10 @@ def initialize_model_and_optimizer(
|
||||
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()
|
||||
|
||||
@@ -1,21 +1,20 @@
|
||||
import torch
|
||||
from megatron.core.dist_checkpointing.strategies.filesystem_async import (
|
||||
FileSystemWriterAsync,
|
||||
)
|
||||
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
|
||||
# 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")
|
||||
@@ -24,6 +23,5 @@ class ROCmFileSystemWriterAsync(FileSystemWriterAsync):
|
||||
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)
|
||||
|
||||
return FileSystemWriterAsync.preload_tensors(*args, **kwargs)
|
||||
|
||||
@@ -75,6 +75,7 @@ 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")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user