Q-Prefer-D2 / code /Q-Prefer /scripts /prepare_adapter.py
qgfvadfuvads's picture
Upload Q-Prefer training and inference code
aa7758f verified
Raw
History Blame Contribute Delete
6.14 kB
#!/usr/bin/env python3
"""Create a portable Q-Prefer adapter directory from the validated D2 run."""
from __future__ import annotations
import argparse
import hashlib
import json
import shutil
from pathlib import Path
from transformers import AutoTokenizer
from qprefer_reward.constants import (
BASE_MODEL_ID,
BASE_MODEL_REVISION,
EXPECTED_SPECIAL_TOKEN_IDS,
PUBLISHED_ADAPTER_SHA256,
SPECIAL_TOKENS,
)
def sha256(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
for chunk in iter(lambda: handle.read(1024 * 1024), b""):
digest.update(chunk)
return digest.hexdigest()
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument(
"--source",
type=Path,
required=True,
help="Validated D2 run or checkpoint directory",
)
parser.add_argument("--output", type=Path, required=True, help="Portable output directory")
parser.add_argument("--base-model", default=BASE_MODEL_ID)
parser.add_argument(
"--base-model-source",
default=None,
help=(
"Optional local snapshot used for offline export; the public "
"--base-model is still written to metadata"
),
)
parser.add_argument("--base-revision", default=BASE_MODEL_REVISION)
parser.add_argument(
"--model-name",
default="Q-Prefer-D2-broad-near",
help="Model name written to artifact_manifest.json",
)
parser.add_argument(
"--special-embeddings-source",
type=Path,
help=(
"Exact special_token_embeddings.safetensors to use when the source "
"is the historical D2 checkpoint, which did not save this file"
),
)
parser.add_argument("--allow-unknown-checkpoint", action="store_true")
parser.add_argument("--overwrite", action="store_true")
return parser.parse_args()
def ensure_output(path: Path, overwrite: bool) -> None:
if path.exists() and any(path.iterdir()) and not overwrite:
raise FileExistsError(f"{path} is not empty; pass --overwrite to replace release files")
path.mkdir(parents=True, exist_ok=True)
def prepare_tokenizer(args: argparse.Namespace, output: Path) -> tuple[object, tuple[int, ...]]:
model_source = args.base_model_source or args.base_model
revision_kwargs = {} if Path(model_source).exists() else {"revision": args.base_revision}
tokenizer = AutoTokenizer.from_pretrained(model_source, use_fast=False, **revision_kwargs)
tokenizer.add_special_tokens({"additional_special_tokens": list(SPECIAL_TOKENS)})
token_ids = tuple(tokenizer.convert_tokens_to_ids(list(SPECIAL_TOKENS)))
if token_ids != EXPECTED_SPECIAL_TOKEN_IDS:
raise RuntimeError(
"unexpected special token ids: "
f"expected={EXPECTED_SPECIAL_TOKEN_IDS}, observed={token_ids}"
)
tokenizer.save_pretrained(output)
return tokenizer, token_ids
def main() -> None:
args = parse_args()
source = args.source.expanduser().resolve()
output = args.output.expanduser().resolve()
adapter_weights = source / "adapter_model.safetensors"
adapter_config = source / "adapter_config.json"
if not adapter_weights.is_file() or not adapter_config.is_file():
raise FileNotFoundError(
"source must contain adapter_model.safetensors and adapter_config.json"
)
observed_sha = sha256(adapter_weights)
if observed_sha != PUBLISHED_ADAPTER_SHA256 and not args.allow_unknown_checkpoint:
raise RuntimeError(
"source is not the published D2 adapter: "
f"expected sha256={PUBLISHED_ADAPTER_SHA256}, observed={observed_sha}. "
"Pass --allow-unknown-checkpoint only for an intentional new model."
)
ensure_output(output, args.overwrite)
shutil.copy2(adapter_weights, output / adapter_weights.name)
model_card = Path(__file__).resolve().parents[1] / "MODEL_CARD.md"
if model_card.is_file():
shutil.copy2(model_card, output / "README.md")
config = json.loads(adapter_config.read_text())
config["base_model_name_or_path"] = args.base_model
(output / "adapter_config.json").write_text(json.dumps(config, indent=2) + "\n")
_, token_ids = prepare_tokenizer(args, output)
source_embeddings = source / "special_token_embeddings.safetensors"
if not source_embeddings.is_file() and args.special_embeddings_source:
source_embeddings = args.special_embeddings_source.expanduser().resolve()
if source_embeddings.is_file():
# New training runs save the exact rows used during optimization. Never
# replace those rows with a fresh base-model resize.
shutil.copy2(source_embeddings, output / "special_token_embeddings.safetensors")
else:
raise FileNotFoundError(
"special_token_embeddings.safetensors is missing. New runs produced by this "
"repository save it automatically. For the historical D2 checkpoint, pass "
"--special-embeddings-source pointing to the exact file from the published artifact; "
"freshly resizing the base model is not an exact replacement."
)
files = {}
for path in sorted(output.iterdir()):
if path.is_file() and path.name != "artifact_manifest.json":
files[path.name] = {"bytes": path.stat().st_size, "sha256": sha256(path)}
manifest = {
"format_version": 1,
"model": args.model_name,
"base_model": args.base_model,
"base_revision": args.base_revision,
"special_tokens": list(SPECIAL_TOKENS),
"special_token_ids": list(token_ids),
"supported_dimensions": ["visual_quality", "text_alignment"],
"unsupported_dimensions": ["motion_quality"],
"files": files,
}
(output / "artifact_manifest.json").write_text(json.dumps(manifest, indent=2) + "\n")
print(f"Prepared portable Q-Prefer adapter at {output}")
print(f"adapter sha256: {observed_sha}")
if __name__ == "__main__":
main()