Enable inference serving && fix train stability (#59)
* enable serving * lint * update * update
This commit is contained in:
@@ -0,0 +1,238 @@
|
||||
#!/usr/bin/env python
|
||||
"""
|
||||
Custom script to merge FSDP sharded checkpoints with compatibility handling.
|
||||
Works around the StorageMeta compatibility issue between PyTorch versions.
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
import torch
|
||||
from pathlib import Path
|
||||
from typing import Dict
|
||||
from safetensors.torch import save_file
|
||||
|
||||
|
||||
def patch_metadata_loader():
|
||||
"""Patch the metadata loader to handle missing StorageMeta class."""
|
||||
import torch.distributed.checkpoint.metadata as metadata_module
|
||||
|
||||
# Create a dummy StorageMeta class if it doesn't exist
|
||||
if not hasattr(metadata_module, "StorageMeta"):
|
||||
print("[INFO] Creating StorageMeta compatibility shim")
|
||||
|
||||
class StorageMeta:
|
||||
"""Compatibility shim for old StorageMeta class."""
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
# Store all args as attributes
|
||||
self.args = args
|
||||
self.kwargs = kwargs
|
||||
|
||||
# Inject the class into the module
|
||||
metadata_module.StorageMeta = StorageMeta
|
||||
|
||||
# Also make it available for unpickling
|
||||
sys.modules["torch.distributed.checkpoint.metadata"].StorageMeta = StorageMeta
|
||||
|
||||
|
||||
def load_sharded_checkpoint(checkpoint_dir: str) -> Dict[str, torch.Tensor]:
|
||||
"""
|
||||
Load a sharded FSDP checkpoint by manually reading all shard files.
|
||||
|
||||
Args:
|
||||
checkpoint_dir: Path to directory containing .distcp files
|
||||
|
||||
Returns:
|
||||
Dictionary of merged model state
|
||||
"""
|
||||
import torch.distributed.checkpoint as dist_cp
|
||||
import torch.distributed.checkpoint.format_utils as dist_cp_format_utils
|
||||
|
||||
print(f"[INFO] Loading checkpoint from {checkpoint_dir}")
|
||||
|
||||
# Apply the compatibility patch
|
||||
patch_metadata_loader()
|
||||
|
||||
# Try to load using the standard approach
|
||||
try:
|
||||
state_dict = {}
|
||||
storage_reader = dist_cp.FileSystemReader(checkpoint_dir)
|
||||
|
||||
dist_cp_format_utils._load_state_dict(
|
||||
state_dict,
|
||||
storage_reader=storage_reader,
|
||||
planner=dist_cp_format_utils._EmptyStateDictLoadPlanner(),
|
||||
no_dist=True,
|
||||
)
|
||||
|
||||
print(f"[INFO] Successfully loaded state dict with {len(state_dict)} keys")
|
||||
return state_dict
|
||||
|
||||
except AttributeError as e:
|
||||
if "StorageMeta" in str(e):
|
||||
print(f"[ERROR] StorageMeta compatibility issue: {e}")
|
||||
print("[INFO] Attempting alternative loading method...")
|
||||
return load_checkpoint_alternative(checkpoint_dir)
|
||||
else:
|
||||
raise
|
||||
|
||||
|
||||
def load_checkpoint_alternative(checkpoint_dir: str) -> Dict[str, torch.Tensor]:
|
||||
"""
|
||||
Alternative method to load checkpoint by directly reading shard files.
|
||||
|
||||
Args:
|
||||
checkpoint_dir: Path to directory containing .distcp files
|
||||
|
||||
Returns:
|
||||
Dictionary of merged model state
|
||||
"""
|
||||
checkpoint_path = Path(checkpoint_dir)
|
||||
|
||||
# Find all shard files
|
||||
shard_files = sorted(checkpoint_path.glob("*.distcp"))
|
||||
|
||||
if not shard_files:
|
||||
raise FileNotFoundError(f"No .distcp files found in {checkpoint_dir}")
|
||||
|
||||
print(f"[INFO] Found {len(shard_files)} shard files")
|
||||
|
||||
# Load all shards
|
||||
merged_state = {}
|
||||
|
||||
for shard_file in shard_files:
|
||||
print(f"[INFO] Loading shard: {shard_file.name}")
|
||||
try:
|
||||
shard_data = torch.load(shard_file, map_location="cpu")
|
||||
|
||||
# Merge the shard into the state dict
|
||||
if isinstance(shard_data, dict):
|
||||
for key, value in shard_data.items():
|
||||
if isinstance(value, torch.Tensor):
|
||||
if key in merged_state:
|
||||
# Handle duplicates - concatenate or overwrite based on shape
|
||||
print(f"[WARNING] Duplicate key found: {key}")
|
||||
merged_state[key] = value
|
||||
elif isinstance(value, dict):
|
||||
# Nested dict structure
|
||||
for subkey, subvalue in value.items():
|
||||
full_key = f"{key}.{subkey}" if key else subkey
|
||||
if isinstance(subvalue, torch.Tensor):
|
||||
merged_state[full_key] = subvalue
|
||||
|
||||
except Exception as e:
|
||||
print(f"[WARNING] Failed to load shard {shard_file.name}: {e}")
|
||||
continue
|
||||
|
||||
if not merged_state:
|
||||
raise RuntimeError("Failed to load any checkpoint data from shards")
|
||||
|
||||
print(f"[INFO] Loaded {len(merged_state)} tensors from shards")
|
||||
return merged_state
|
||||
|
||||
|
||||
def save_merged_checkpoint(
|
||||
state_dict: Dict[str, torch.Tensor],
|
||||
output_path: str,
|
||||
safe_serialization: bool = True,
|
||||
):
|
||||
"""
|
||||
Save the merged checkpoint to disk.
|
||||
|
||||
Args:
|
||||
state_dict: Model state dictionary
|
||||
output_path: Directory to save the merged checkpoint
|
||||
safe_serialization: If True, save as .safetensors, else as .bin
|
||||
"""
|
||||
output_dir = Path(output_path)
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Handle nested state dict structure (e.g., {model: {...}})
|
||||
if len(state_dict.keys()) == 1 and all(
|
||||
isinstance(v, dict) for v in state_dict.values()
|
||||
):
|
||||
print("[INFO] Unwrapping nested state dict")
|
||||
state_dict = state_dict[list(state_dict.keys())[0]]
|
||||
|
||||
# Prepare tensors for saving
|
||||
save_dict = {}
|
||||
for key, value in state_dict.items():
|
||||
if isinstance(value, torch.Tensor):
|
||||
# Convert to CPU and contiguous
|
||||
save_dict[key] = value.cpu().contiguous()
|
||||
else:
|
||||
print(f"[WARNING] Skipping non-tensor key: {key} (type: {type(value)})")
|
||||
|
||||
if safe_serialization:
|
||||
output_file = output_dir / "model.safetensors"
|
||||
print(f"[INFO] Saving merged checkpoint to {output_file}")
|
||||
save_file(save_dict, output_file)
|
||||
else:
|
||||
output_file = output_dir / "pytorch_model.bin"
|
||||
print(f"[INFO] Saving merged checkpoint to {output_file}")
|
||||
torch.save(save_dict, output_file)
|
||||
|
||||
print("[SUCCESS] Checkpoint saved successfully!")
|
||||
print(f"[INFO] Saved {len(save_dict)} tensors")
|
||||
|
||||
# Print size info
|
||||
total_params = sum(v.numel() for v in save_dict.values())
|
||||
total_size_gb = sum(v.numel() * v.element_size() for v in save_dict.values()) / (
|
||||
1024**3
|
||||
)
|
||||
print(f"[INFO] Total parameters: {total_params:,}")
|
||||
print(f"[INFO] Total size: {total_size_gb:.2f} GB")
|
||||
|
||||
return output_file
|
||||
|
||||
|
||||
def main():
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Merge FSDP sharded checkpoints with compatibility handling"
|
||||
)
|
||||
parser.add_argument(
|
||||
"checkpoint_dir",
|
||||
type=str,
|
||||
help="Directory containing sharded FSDP checkpoint files (*.distcp)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"output_path", type=str, help="Output directory for merged checkpoint"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--unsafe-serialization",
|
||||
action="store_true",
|
||||
help="Save as .bin instead of .safetensors",
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
# Validate input
|
||||
if not os.path.exists(args.checkpoint_dir):
|
||||
print(f"[ERROR] Checkpoint directory not found: {args.checkpoint_dir}")
|
||||
sys.exit(1)
|
||||
|
||||
try:
|
||||
# Load the sharded checkpoint
|
||||
state_dict = load_sharded_checkpoint(args.checkpoint_dir)
|
||||
|
||||
# Save the merged checkpoint
|
||||
safe_serialization = not args.unsafe_serialization
|
||||
output_file = save_merged_checkpoint(
|
||||
state_dict, args.output_path, safe_serialization
|
||||
)
|
||||
|
||||
print("\n[COMPLETE] Checkpoint merging successful!")
|
||||
print(f"[COMPLETE] Output: {output_file}")
|
||||
|
||||
except Exception as e:
|
||||
print(f"\n[ERROR] Failed to merge checkpoint: {e}")
|
||||
import traceback
|
||||
|
||||
traceback.print_exc()
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user