Skip to content
Merged
Show file tree
Hide file tree
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
179 changes: 179 additions & 0 deletions areal/experimental/engine/archon_checkpoint.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,179 @@
from __future__ import annotations

import os
import shutil
from typing import TYPE_CHECKING

import torch
import torch.distributed as dist
import torch.distributed.checkpoint as dcp
from torch.distributed.checkpoint.state_dict import (
StateDictOptions,
get_model_state_dict,
set_model_state_dict,
)

from areal.utils.fsdp.checkpoint import DCPState

if TYPE_CHECKING:
from transformers import AutoProcessor, PreTrainedTokenizerFast

from areal.experimental.engine.archon_engine import ArchonEngine


def save_model_to_hf(
engine: ArchonEngine,
path: str,
tokenizer: PreTrainedTokenizerFast | None,
processor: AutoProcessor | None = None,
) -> None:
"""Save model in HuggingFace format using DCP infrastructure."""
from torch.distributed.checkpoint import HuggingFaceStorageWriter

if engine.model is None:
raise RuntimeError("Model not initialized")
if engine.state_dict_adapter is None:
raise RuntimeError("state_dict_adapter is required for HF format")

engine.logger.info(f"Saving HF checkpoint to {path}")
os.makedirs(path, exist_ok=True)

# Get distributed state dict
options = StateDictOptions(full_state_dict=False, cpu_offload=True)
state_dict = get_model_state_dict(engine.model, options=options)

# Convert to HF format using adapter
hf_state_dict = engine.state_dict_adapter.to_hf(state_dict)

fqn_to_index_mapping = engine.state_dict_adapter.fqn_to_index_mapping

# NOTE: HuggingFaceStorageWriter always creates a sharded/ subdirectory when
# save_distributed=True. With enable_consolidation=True, it saves shards to
# path/sharded/, then consolidates to path/. The sharded/ directory is NOT
# automatically cleaned up by PyTorch, so we must remove it manually.
sharded_dir = os.path.join(path, "sharded")

if fqn_to_index_mapping:
# Multi-file output: save to sharded/, then consolidate with all ranks
from torch.distributed.checkpoint._consolidate_hf_safetensors import (
consolidate_safetensors_files_on_every_rank,
)

hf_writer = HuggingFaceStorageWriter(
path=sharded_dir,
save_distributed=True,
fqn_to_index_mapping=fqn_to_index_mapping,
enable_consolidation=False,
)
dcp.save(hf_state_dict, storage_writer=hf_writer)

# NOTE: consolidate_safetensors_files_on_every_rank() has internal barrier
consolidate_safetensors_files_on_every_rank(
input_dir=sharded_dir,
output_dir=path,
fqn_to_index_mapping=fqn_to_index_mapping,
num_threads=8,
)
Comment on lines +62 to +76

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

What about the performance?

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Our benchmark on a 12.55 GB model across 8 GPUs shows that DCP provides a modest 1.10x speedup (11.18s vs 12.27s), but the real advantage is memory efficiency: DCP uses 23x less CPU memory (0.45 GB vs 10.52 GB). The baseline approach with full_state_dict=True gathers the entire model to each rank before saving, consuming ~10.5 GB extra CPU memory per rank—this becomes the bottleneck for larger models where OOM is inevitable. DCP keeps data sharded throughout the save process, using only ~1.57 GB per rank. The time breakdown reveals that DCP's get_state_dict is 7x faster (0.39s vs 2.72s) by avoiding all-gather communication, though the two-phase write (DCP save + consolidate) takes longer than single-rank sequential write. In summary, DCP's primary value is enabling checkpoint saves for large models that would otherwise OOM, with a secondary benefit of slightly better throughput at scale.

Benchmarking: https://gist.github.com/rchardx/bc558b7bf7cab969e4075a71b17ed1a9

torchrun --nproc_per_node=8benchmark_hf_checkpoint_save.py --model-size large --warmup
Creating FSDP2 model: large
  World size: 8
  Estimated size: 12.55 GB
  Per-GPU shard: 1.57 GB
Model created and wrapped with FSDP2
  GPU memory after model init: 1.38 GB

Running warmup...

Running BASELINE benchmark...
Running DCP benchmark...

================================================================================
Model: large
  Layers: 32, Hidden: 4096
  Estimated size: 12.55 GB
  World size: 8 GPUs
  Per-GPU shard size: 1.57 GB
================================================================================

BASELINE (full_state_dict=True, rank 0 saves):
  Time:
    All-gather: 2.72s
    Save (rank 0): 9.54s
    Total: 12.27s
    Throughput: 1.02 GB/s
  Memory (rank 0):
    GPU before: 1.38 GB
    GPU peak (after gather): 1.63 GB
    GPU delta: +0.24 GB
    CPU before: 2.10 GB
    CPU after gather: 12.62 GB
    CPU delta: +10.52 GB

DCP (full_state_dict=False, parallel save + consolidate):
  Time:
    Get state dict: 0.39s
    DCP save: 7.51s
    Consolidate: 3.28s
    Total: 11.18s
    Throughput: 1.12 GB/s
  Memory (rank 0):
    GPU before: 1.38 GB
    GPU peak (after get): 1.38 GB
    GPU delta: +0.00 GB
    CPU before: 2.10 GB
    CPU after get: 2.55 GB
    CPU delta: +0.45 GB

--------------------------------------------------------------------------------
COMPARISON:
--------------------------------------------------------------------------------
  Time speedup (DCP vs Baseline): 1.10x
  CPU memory ratio (Baseline / DCP): 23.29x
    Baseline needs 10.52 GB extra CPU memory
    DCP needs 0.45 GB extra CPU memory

  Theoretical analysis:
    Model size: 12.55 GB
    Baseline gathers full model to each rank -> ~12.55 GB extra memory
    DCP keeps sharded data -> ~1.57 GB per rank
================================================================================

else:
# Single-file output: auto consolidation
hf_writer = HuggingFaceStorageWriter(
path=path,
save_distributed=True,
enable_consolidation=True,
)
dcp.save(hf_state_dict, storage_writer=hf_writer)

# Clean up sharded/ directory after consolidation
if dist.get_rank() == 0 and os.path.exists(sharded_dir):
shutil.rmtree(sharded_dir)

dist.barrier(group=engine.cpu_group)

if dist.get_rank() == 0:
engine.model_config.save_pretrained(path)
if tokenizer is not None:
tokenizer.save_pretrained(path)
if processor is not None:
processor.save_pretrained(path)
dist.barrier(group=engine.cpu_group)


def load_model_from_hf(engine: ArchonEngine, path: str) -> None:
"""Load model from HuggingFace format using DCP infrastructure."""
if engine.model is None:
raise RuntimeError("Model not initialized")
if engine.state_dict_adapter is None:
raise RuntimeError("state_dict_adapter is required for HF format")

engine.logger.info(f"Loading HF checkpoint from {path}")

# Get model state dict structure (distributed)
options = StateDictOptions(full_state_dict=False, cpu_offload=True)
state_dict = get_model_state_dict(engine.model, options=options)

# Convert to HF format to match checkpoint keys
hf_state_dict = engine.state_dict_adapter.to_hf(state_dict)

# Load using DCP with HuggingFaceStorageReader
hf_reader = engine.state_dict_adapter.get_hf_storage_reader(path)
dcp.load(hf_state_dict, storage_reader=hf_reader)

# Convert back to Archon format
archon_state_dict = engine.state_dict_adapter.from_hf(hf_state_dict)

# Load into FSDP model (same as DCPState.load_state_dict)
set_model_state_dict(
engine.model,
model_state_dict=archon_state_dict,
options=StateDictOptions(strict=False),
)

dist.barrier(group=engine.cpu_group)


def save_to_dcp(engine: ArchonEngine, path: str, with_optim: bool) -> None:
"""Save model (and optionally optimizer) using DCP format."""
if engine.model is None:
raise RuntimeError("Model not initialized")

os.makedirs(path, exist_ok=True)

dcp_state = DCPState(engine.model, engine.optimizer if with_optim else None)
state_dict = {"dcp": dcp_state}
dcp.save(state_dict, checkpoint_id=path)


def load_from_dcp(engine: ArchonEngine, path: str, with_optim: bool) -> None:
"""Load model (and optionally optimizer) from DCP format."""
if engine.model is None:
raise RuntimeError("Model not initialized")

dcp_state = DCPState(engine.model, engine.optimizer if with_optim else None)
state_dict = {"dcp": dcp_state}
dcp.load(state_dict=state_dict, checkpoint_id=path)


def save_optimizer_state(engine: ArchonEngine, path: str) -> None:
"""Save optimizer state to disk (sharded by rank)."""
assert engine.optimizer is not None
assert dist.is_initialized()
rank = dist.get_rank()
shard_path = os.path.join(
path, f"optim_world_size_{engine.world_size}_rank_{rank}.pt"
)
state_dict = engine.optimizer.state_dict()
torch.save(state_dict, shard_path)
dist.barrier(group=engine.cpu_group)


def load_optimizer_state(engine: ArchonEngine, path: str) -> None:
"""Load optimizer state from disk (sharded by rank)."""
assert engine.optimizer is not None
assert dist.is_initialized()
rank = dist.get_rank()
shard_path = os.path.join(
path, f"optim_world_size_{engine.world_size}_rank_{rank}.pt"
)
optimizer_state_dict = torch.load(shard_path, weights_only=False)
engine.optimizer.load_state_dict(optimizer_state_dict)
dist.barrier(group=engine.cpu_group)
124 changes: 19 additions & 105 deletions areal/experimental/engine/archon_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,13 +12,7 @@

import torch
import torch.distributed as dist
import torch.distributed.checkpoint as dcp
from safetensors.torch import save_file
from torch import nn
from torch.distributed.checkpoint.state_dict import (
StateDictOptions,
get_model_state_dict,
)
from torch.distributed.device_mesh import DeviceMesh
from torch.distributed.tensor import DTensor
from torchdata.stateful_dataloader import StatefulDataLoader
Expand Down Expand Up @@ -48,6 +42,14 @@
compute_total_loss_weight,
reorder_and_pad_outputs,
)
from areal.experimental.engine.archon_checkpoint import (
load_from_dcp,
load_model_from_hf,
load_optimizer_state,
save_model_to_hf,
save_optimizer_state,
save_to_dcp,
)
from areal.experimental.models.archon import (
ArchonParallelDims,
BaseStateDictAdapter,
Expand Down Expand Up @@ -76,8 +78,7 @@
unsqueeze_mb_list,
)
from areal.utils.distributed import init_custom_process_group, patch_dist_group_timeout
from areal.utils.fsdp import fsdp2_load_full_state_dict, get_cosine_schedule_with_warmup
from areal.utils.fsdp.checkpoint import DCPState
from areal.utils.fsdp import get_cosine_schedule_with_warmup
from areal.utils.fsdp.grad import fsdp2_clip_grad_norm
from areal.utils.functional import gather_logprobs, gather_logprobs_entropy
from areal.utils.hf_utils import load_hf_tokenizer
Expand Down Expand Up @@ -562,26 +563,26 @@ def update_weights(self, meta: WeightUpdateMeta):
def save(self, meta: SaveLoadMeta):
"""Save model in HuggingFace or DCP format."""
if meta.weight_format == "hf":
self._save_model_to_hf(meta.path, meta.tokenizer, meta.processor)
save_model_to_hf(self, meta.path, meta.tokenizer, meta.processor)
elif meta.weight_format == "dcp":
self._save_to_dcp(meta.path, meta.with_optim)
save_to_dcp(self, meta.path, meta.with_optim)
else:
raise ValueError(f"Unknown weight format {meta.weight_format}.")

if meta.with_optim and meta.weight_format == "hf":
self._save_optimizer_state(meta.path)
save_optimizer_state(self, meta.path)

def load(self, meta: SaveLoadMeta):
"""Load model from HuggingFace or DCP format."""
if meta.weight_format == "hf":
self._load_model_from_hf(meta.path)
load_model_from_hf(self, meta.path)
elif meta.weight_format == "dcp":
self._load_from_dcp(meta.path, meta.with_optim)
load_from_dcp(self, meta.path, meta.with_optim)
else:
raise ValueError(f"Unknown weight format {meta.weight_format}.")

if meta.with_optim and meta.weight_format == "hf":
self._load_optimizer_state(meta.path)
load_optimizer_state(self, meta.path)

def offload(self) -> None:
"""Offload model memory to CPU using torch_memory_saver."""
Expand Down Expand Up @@ -642,7 +643,9 @@ def _validate_model_type(self) -> None:
)

def _create_state_dict_adapter(self) -> BaseStateDictAdapter | None:
return self.spec.state_dict_adapter_class(self.model_config)
return self.spec.state_dict_adapter_class(
self.model_config, hf_assets_path=self.config.path
)

def _get_model_name_parameters(self) -> Iterator[tuple[str, nn.Parameter]]:
return self.model.named_parameters()
Expand Down Expand Up @@ -790,7 +793,7 @@ def _update_weights_from_disk(self, meta: WeightUpdateMeta):
fut = self.rollout_engine.update_weights_from_disk(meta)

assert meta.path is not None
self._save_model_to_hf(meta.path, self.tokenizer, None)
save_model_to_hf(self, meta.path, self.tokenizer, None)

if dist.get_rank() == 0:
update_name = names.update_weights_from_disk(
Expand All @@ -807,95 +810,6 @@ def _update_weights_from_disk(self, meta: WeightUpdateMeta):
current_platform.synchronize()
dist.barrier(group=self.cpu_group)

def _save_model_to_hf(
self,
path: str,
tokenizer: PreTrainedTokenizerFast | None,
processor=None,
):
"""Save model in HuggingFace format."""
if self.model is None:
raise RuntimeError("Model not initialized")
os.makedirs(path, exist_ok=True)

options = StateDictOptions(full_state_dict=True, cpu_offload=True)
state_dict = get_model_state_dict(self.model, options=options)
if self.state_dict_adapter is not None:
state_dict = self.state_dict_adapter.to_hf(state_dict)

if dist.get_rank() == 0:
os.makedirs(path, exist_ok=True)
model_path = os.path.join(path, "model.safetensors")
try:
save_file(state_dict, model_path)
except ImportError:
model_path = os.path.join(path, "pytorch_model.bin")
torch.save(state_dict, model_path)

self.model_config.save_pretrained(path)
if tokenizer is not None:
tokenizer.save_pretrained(path)
if processor is not None:
processor.save_pretrained(path)
dist.barrier(group=self.cpu_group)

def _load_model_from_hf(self, path: str):
"""Load model from HuggingFace format."""
if dist.get_rank() == 0:
full_state = get_state_dict_from_repo_id_or_path(path)
if self.state_dict_adapter is not None:
full_state = self.state_dict_adapter.from_hf(full_state)
else:
full_state = {}

cpu_offload = self.config.archon.offload_params
fsdp2_load_full_state_dict(
self.model,
full_state,
cpu_offload,
tie_word_embeddings=self.model_config.tie_word_embeddings,
)

def _save_to_dcp(self, path: str, with_optim: bool):
if self.model is None:
raise RuntimeError("Model not initialized")

os.makedirs(path, exist_ok=True)

dcp_state = DCPState(self.model, self.optimizer if with_optim else None)
state_dict = {"dcp": dcp_state}
dcp.save(state_dict, checkpoint_id=path)

def _load_from_dcp(self, path: str, with_optim: bool):
if self.model is None:
raise RuntimeError("Model not initialized")

dcp_state = DCPState(self.model, self.optimizer if with_optim else None)
state_dict = {"dcp": dcp_state}
dcp.load(state_dict=state_dict, checkpoint_id=path)

def _save_optimizer_state(self, path: str):
assert self.optimizer is not None
assert dist.is_initialized()
rank = dist.get_rank()
shard_path = os.path.join(
path, f"optim_world_size_{self.world_size}_rank_{rank}.pt"
)
state_dict = self.optimizer.state_dict()
torch.save(state_dict, shard_path)
dist.barrier(group=self.cpu_group)

def _load_optimizer_state(self, path: str):
assert self.optimizer is not None
assert dist.is_initialized()
rank = dist.get_rank()
shard_path = os.path.join(
path, f"optim_world_size_{self.world_size}_rank_{rank}.pt"
)
optimizer_state_dict = torch.load(shard_path, weights_only=False)
self.optimizer.load_state_dict(optimizer_state_dict)
dist.barrier(group=self.cpu_group)

def _create_device_model(self):
current_platform.set_device(int(os.environ["LOCAL_RANK"]))
self.device = torch.device(int(os.environ["LOCAL_RANK"]))
Expand Down
6 changes: 6 additions & 0 deletions areal/experimental/models/archon/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,10 @@
get_supported_model_types,
is_supported_model,
)
from areal.experimental.models.archon.moe_weight_converter import (
MoEConversionState,
MoEWeightConverter,
)
from areal.experimental.models.archon.parallel_dims import (
ArchonParallelDims,
)
Expand All @@ -25,6 +29,8 @@
"BaseStateDictAdapter",
"ExpertParallel",
"ExpertTensorParallel",
"MoEConversionState",
"MoEWeightConverter",
"ModelSpec",
"TensorParallel",
"apply_expert_parallel",
Expand Down
Loading