#!/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()