| |
| """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() |
|
|