Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
109 changes: 42 additions & 67 deletions megatron/core/dist_checkpointing/strategies/filesystem_async.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,11 @@
from torch import multiprocessing as mp
from torch.distributed.checkpoint import FileSystemWriter
from torch.distributed.checkpoint.api import WRAPPED_EXCEPTION, _wrap_exception
from torch.distributed.checkpoint.filesystem import DEFAULT_SUFFIX, _StoragePrefix, _write_item
from torch.distributed.checkpoint.filesystem import (
DEFAULT_SUFFIX,
_StoragePrefix,
_write_item,
)
from torch.distributed.checkpoint.metadata import Metadata

try:
Expand Down Expand Up @@ -240,19 +244,12 @@ def write_preloaded_data_multiproc(
global_results_queue: mp.Queue,
) -> None:
"""
Performs saving data to storage with multiple processes.
Performs saving data to storage in the current process.

Starts predefined number of processes and uses 2 queues to make sure the results
are complete:
- local_results_queue - to send the actual results
- count_queue - small queue to mark worker as completed

Using just one queue disallowed proper exception handling.

This method is meant to be run in a forked subprocess.
Triggering GC during execution leads to CUDA errors
(cleaning up tensors owned by the parent process).
To prevent this, we disable the GC explicitly for this function with _disable_gc.
This keeps the public entry point used by async checkpointing but avoids spawning
local writer subprocesses. Forked local writers can read CPU tensors staged in
the parent process and crash inside PyTorch serialization on some platforms.
GC is still disabled during execution because tensors are owned by the caller.

Args:
write_buckets (List[WriteBucket]): write plan
Expand All @@ -262,80 +259,53 @@ def write_preloaded_data_multiproc(
"""
logger = logging.getLogger(__name__)
w_start = time()
write_results_or_exc: Union[dict, Exception] = dict()
ctx = mp.get_context("fork")
local_results_queue = ctx.Queue()
count_queue = ctx.JoinableQueue()
p_list = []
for i, write_bucket in enumerate(write_buckets):
write_results_or_exc: Union[dict, Exception] = {}
for local_proc_idx, write_bucket in enumerate(write_buckets):
local_results_queue = queue.Queue()
count_queue = queue.Queue()
try:
count_queue.put(i)

count_queue.put(local_proc_idx)
kwargs = {
"local_proc_idx": i,
"local_proc_idx": local_proc_idx,
"write_bucket": write_bucket,
"results_queue": local_results_queue,
"count_queue": count_queue,
"use_fsync": True,
}

if use_msc:
import inspect

# Remove the inspect after the test_async_save.py is fixed.
signature = inspect.signature(FileSystemWriterAsync.write_preloaded_data)
if len(signature.parameters) > 6:
kwargs["use_msc"] = use_msc

p_list.append(
ctx.Process(
target=partial(FileSystemWriterAsync.write_preloaded_data, transform_list),
kwargs=kwargs,
FileSystemWriterAsync.write_preloaded_data(transform_list, **kwargs)
local_proc_idx, local_results_or_exc = local_results_queue.get_nowait()
if isinstance(local_results_or_exc, Exception):
err_msg = (
f"Local process {local_proc_idx} encountered"
f" an error: {local_results_or_exc}"
)
logger.error(err_msg)
write_results_or_exc = local_results_or_exc
break
assert isinstance(local_results_or_exc, list), type(local_results_or_exc)
write_results_or_exc[local_proc_idx] = local_results_or_exc
except queue.Empty:
write_results_or_exc = RuntimeError(
"Unexpected empty `local_results_queue`"
f" for bucket {local_proc_idx}/{len(write_buckets)}"
)
break
except Exception as e:
err_msg = f"An error is caught while a proc {i} is created, error: {e}"
logger.error(err_msg)
write_results_or_exc = RuntimeError(err_msg)

if not isinstance(write_results_or_exc, Exception):
for p in p_list:
p.start()

logger.debug("FileSystemWriterAsync: collecting worker results...")

# To make sure all nodes are completed
count_queue.join()
# At this point, all workers completed, so the queue should have exactly
# `len(write_buckets)` items
for proc_idx in range(len(write_buckets)):
try:
local_proc_idx, local_results_or_exc = local_results_queue.get()
except queue.Empty:
write_results_or_exc = RuntimeError(
"Unexpected empty `local_results_queue`"
f" (got only {proc_idx}/{len(write_buckets)} items)"
)
break
else:
if isinstance(local_results_or_exc, Exception):
err_msg = (
f"Local process {local_proc_idx} encountered"
f" an error: {local_results_or_exc}"
)
logger.error(err_msg)
write_results_or_exc = local_results_or_exc
break
assert isinstance(local_results_or_exc, list), type(local_results_or_exc)
write_results_or_exc[local_proc_idx] = local_results_or_exc
p_list[local_proc_idx].join()

logger.debug("FileSystemWriterAsync: collected worker results successfully")
logger.exception("Checkpoint local write bucket %s failed", local_proc_idx)
write_results_or_exc = e
break

global_results_queue.put(write_results_or_exc)

w_end = time()
logger.debug(f"{w_end}, rank: {rank}, write(sync,parallel): {w_end - w_start}")
logger.debug(f"{w_end}, rank: {rank}, write(sync,sequential): {w_end - w_start}")

@staticmethod
@_disable_gc()
Expand Down Expand Up @@ -392,7 +362,12 @@ def write_preloaded_data(
assert tensor.is_cpu
local_results.append(
_write_item(
*transform_list, stream, tensor, write_item, storage_key, **extra_kwargs
*transform_list,
stream,
tensor,
write_item,
storage_key,
**extra_kwargs,
)
)

Expand Down