hep-posttraining / dataset /scripts /export_verl_lora_adapter.py
ho22joshua's picture
Upload folder using huggingface_hub
e8f2c80 verified
Raw
History Blame Contribute Delete
4.29 kB
#!/usr/bin/env python3
"""Export a VERL FSDP LoRA checkpoint as a Hugging Face PEFT adapter."""
from __future__ import annotations
import argparse
import json
import shutil
from pathlib import Path
import torch
from peft import LoraConfig
from safetensors.torch import save_file
SIDECAR_FILES = [
"config.json",
"generation_config.json",
"tokenizer.json",
"tokenizer_config.json",
"special_tokens_map.json",
"added_tokens.json",
"vocab.json",
"merges.txt",
"chat_template.jinja",
]
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description=(
"Extract LoRA tensors from a VERL FSDP checkpoint without reconstructing "
"the full base model."
)
)
parser.add_argument("--checkpoint-dir", required=True, help="VERL global_step_* checkpoint directory.")
parser.add_argument("--output-dir", required=True, help="Output PEFT adapter directory.")
parser.add_argument("--base-model", default="Qwen/Qwen2.5-7B-Instruct")
parser.add_argument("--copy-tokenizer", action="store_true", help="Copy tokenizer/config sidecars if present.")
return parser.parse_args()
def load_json(path: Path) -> dict[str, object]:
with path.open(encoding="utf-8") as handle:
return json.load(handle)
def peft_key_from_verl_key(key: str) -> str:
return key.replace(".default.weight", ".weight")
def infer_target_module(key: str) -> str:
return key.split(".lora_", maxsplit=1)[0].rsplit(".", maxsplit=1)[-1]
def main() -> None:
args = parse_args()
checkpoint_dir = Path(args.checkpoint_dir)
output_dir = Path(args.output_dir)
output_dir.mkdir(parents=True, exist_ok=True)
fsdp_config = load_json(checkpoint_dir / "fsdp_config.json")
world_size = int(fsdp_config["world_size"])
lora_meta_path = checkpoint_dir / "lora_train_meta.json"
lora_meta = load_json(lora_meta_path) if lora_meta_path.exists() else {}
tensor_parts: dict[str, list[torch.Tensor]] = {}
target_modules: set[str] = set()
for rank in range(world_size):
shard_path = checkpoint_dir / f"model_world_size_{world_size}_rank_{rank}.pt"
state_dict = torch.load(shard_path, map_location="cpu", weights_only=False)
for key, value in state_dict.items():
if "lora_" not in key:
continue
tensor = value._local_tensor if hasattr(value, "_local_tensor") else value
peft_key = peft_key_from_verl_key(key)
tensor_parts.setdefault(peft_key, []).append(tensor.detach().cpu().contiguous().bfloat16())
target_modules.add(infer_target_module(key))
del state_dict
print(f"loaded LoRA tensors from rank {rank}")
if not tensor_parts:
raise SystemExit(f"No LoRA tensors found in {checkpoint_dir}")
adapter_state = {key: torch.cat(parts, dim=0).contiguous() for key, parts in sorted(tensor_parts.items())}
save_file(adapter_state, output_dir / "adapter_model.safetensors")
lora_config = LoraConfig(
r=int(lora_meta.get("r", 16)),
lora_alpha=int(lora_meta.get("lora_alpha", lora_meta.get("r", 16))),
target_modules=sorted(target_modules),
lora_dropout=0.0,
bias="none",
task_type=str(lora_meta.get("task_type", "CAUSAL_LM")),
)
lora_config.base_model_name_or_path = args.base_model
lora_config.save_pretrained(output_dir)
if args.copy_tokenizer:
hf_dir = checkpoint_dir / "huggingface"
for name in SIDECAR_FILES:
src = hf_dir / name
if src.exists():
shutil.copy2(src, output_dir / name)
metadata = {
"adapter_tensors": len(adapter_state),
"base_model": args.base_model,
"checkpoint_dir": str(checkpoint_dir),
"lora_alpha": int(lora_meta.get("lora_alpha", lora_meta.get("r", 16))),
"rank": int(lora_meta.get("r", 16)),
"target_modules": sorted(target_modules),
"world_size": world_size,
}
(output_dir / "export_metadata.json").write_text(json.dumps(metadata, indent=2, sort_keys=True) + "\n")
print(json.dumps(metadata, indent=2, sort_keys=True))
print(f"Wrote PEFT adapter to {output_dir}")
if __name__ == "__main__":
main()