File size: 4,213 Bytes
2abcc30
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
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