[Hardware] AMD - pre-commit to fix the formate in commit 4f6b95f (#302)

This commit is contained in:
Ethan (Yusheng) Su
2025-12-05 11:57:54 -08:00
committed by GitHub
parent 4f6b95f02e
commit 6e70fb54ac
3 changed files with 9 additions and 9 deletions
+2 -1
View File
@@ -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()
+6 -8
View File
@@ -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)
+1
View File
@@ -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")