from __future__ import annotations import json import shutil from pathlib import Path from .constants import ( ALLOWED_EXPORT_EXTRAS, DEFAULT_ADAPTER_DIR, DEFAULT_EXPORT_DIR, GENESIS_DIR, METADATA_FILES, STRIP_EXPORT_NAMES, ) def export_merged( *, adapter_dir: Path = DEFAULT_ADAPTER_DIR, genesis_dir: Path = GENESIS_DIR, out_dir: Path = DEFAULT_EXPORT_DIR, base_dir: Path | None = None, ) -> Path: """Merge LoRA (if present) and copy genesis metadata byte-for-byte.""" adapter_dir = Path(adapter_dir) genesis_dir = Path(genesis_dir) out_dir = Path(out_dir) base_dir = Path(base_dir or genesis_dir) out_dir.mkdir(parents=True, exist_ok=True) if (adapter_dir / "adapter_config.json").is_file(): _merge_adapter(base_dir, adapter_dir, out_dir) elif adapter_dir.resolve() != out_dir.resolve() and any(adapter_dir.glob("*.safetensors")): _copy_weights(adapter_dir, out_dir) elif not any(out_dir.glob("*.safetensors")): raise FileNotFoundError(f"no adapter or safetensors under {adapter_dir}") _copy_genesis_metadata(genesis_dir, out_dir) stripped = _strip_disallowed(out_dir) report = { "out_dir": str(out_dir), "adapter_dir": str(adapter_dir), "genesis_dir": str(genesis_dir), "copied_metadata": list(METADATA_FILES), "stripped": stripped, "safetensors": sorted(p.name for p in out_dir.glob("*.safetensors")), } (out_dir / "export-report.json").write_text(json.dumps(report, indent=2) + "\n") # export-report.json is not on the allowlist — remove after writing a copy beside the run. sidecar = out_dir.parent / f"{out_dir.name}-export-report.json" shutil.move(str(out_dir / "export-report.json"), sidecar) print(json.dumps(report, indent=2), flush=True) print(f"export: {out_dir}", flush=True) return out_dir def _merge_adapter(base_dir: Path, adapter_dir: Path, out_dir: Path) -> None: import torch from peft import PeftModel from transformers import AutoModelForCausalLM print(f"merging adapter {adapter_dir} onto {base_dir}", flush=True) model = AutoModelForCausalLM.from_pretrained( str(base_dir), torch_dtype=torch.bfloat16, trust_remote_code=False, device_map="cpu", ) model = PeftModel.from_pretrained(model, str(adapter_dir)) model = model.merge_and_unload() out_dir.mkdir(parents=True, exist_ok=True) model.save_pretrained(str(out_dir), safe_serialization=True) del model def _copy_weights(src: Path, dest: Path) -> None: dest.mkdir(parents=True, exist_ok=True) for path in src.glob("model*.safetensors"): shutil.copy2(path, dest / path.name) index = src / "model.safetensors.index.json" if index.is_file(): shutil.copy2(index, dest / index.name) def _copy_genesis_metadata(genesis_dir: Path, out_dir: Path) -> None: for name in METADATA_FILES: src = genesis_dir / name if not src.is_file(): raise FileNotFoundError(f"genesis metadata missing: {src}") shutil.copy2(src, out_dir / name) for name in ALLOWED_EXPORT_EXTRAS: src = genesis_dir / name if src.is_file() and name != "model.safetensors.index.json": shutil.copy2(src, out_dir / name) def _strip_disallowed(out_dir: Path) -> list[str]: removed: list[str] = [] allowed = set(METADATA_FILES) | set(ALLOWED_EXPORT_EXTRAS) for path in out_dir.iterdir(): if not path.is_file(): if path.is_dir() and path.name.startswith("checkpoint"): shutil.rmtree(path) removed.append(path.name + "/") continue name = path.name if name in allowed or name.endswith(".safetensors"): continue if name.startswith("model-") and name.endswith(".safetensors"): continue if name in STRIP_EXPORT_NAMES or name.endswith(".py") or name.endswith(".bin"): path.unlink() removed.append(name) continue # Anything else also fails the live manifest. path.unlink() removed.append(name) return removed