dougvk's picture
Publish Unlimited-OCR RDNA 4 runtime v0.1.0
f340984 verified
Raw
History Blame Contribute Delete
11.8 kB
from __future__ import annotations
import hashlib
import json
import os
import shutil
import stat
import tempfile
from dataclasses import asdict, dataclass
from fcntl import LOCK_EX, flock
from pathlib import Path
from typing import BinaryIO
from .constants import (
MODEL_DOWNLOAD_SHA256,
MODEL_MANIFEST_FILE,
MODEL_PAYLOAD_SHA256,
MODEL_REPO,
MODEL_REVISION,
MODEL_UPSTREAM_BACKUP_FILE,
MODEL_WEIGHT_BYTES,
MODEL_WEIGHT_FILE,
MODEL_WEIGHT_SHA256,
PATCHED_MODEL_CODE_SHA256,
UPSTREAM_MODEL_CODE_SHA256,
VERIFIED_ROCM_VERSION,
VERIFIED_TORCH_VERSION,
)
from .errors import ModelIntegrityError
@dataclass(frozen=True)
class ModelStatus:
model_dir: str
revision: str
prepared: bool
weight_bytes: int
weight_sha256: str | None
model_code_sha256: str
reused: bool = False
def _regular_file(path: Path) -> bool:
try:
return stat.S_ISREG(path.lstat().st_mode)
except FileNotFoundError:
return False
def sha256_file(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
for chunk in iter(lambda: handle.read(4 * 1024 * 1024), b""):
digest.update(chunk)
return digest.hexdigest()
def sha256_text(text: str) -> str:
return hashlib.sha256(text.encode("utf-8")).hexdigest()
def _replace_exact(text: str, old: str, new: str, expected: int) -> str:
found = text.count(old)
if found != expected:
raise ModelIntegrityError(f"model patch context mismatch: expected {expected} occurrence(s), found {found}")
return text.replace(old, new)
def patch_model_source_text(text: str) -> str:
current_hash = sha256_text(text)
if current_hash == PATCHED_MODEL_CODE_SHA256:
return text
if current_hash != UPSTREAM_MODEL_CODE_SHA256:
raise ModelIntegrityError(
"refusing to patch unknown model code; expected pinned Baidu revision "
f"{MODEL_REVISION}, got SHA-256 {current_hash}"
)
text = _replace_exact(
text,
"from .modeling_deepseekv2 import DeepseekV2Model, DeepseekV2ForCausalLM",
"import ast\n\nfrom .modeling_deepseekv2 import DeepseekV2Model, DeepseekV2ForCausalLM",
1,
)
replacements = (
("cor_list = eval(ref_text[2])", "cor_list = ast.literal_eval(ref_text[2])"),
("lines = eval(outputs)['Line']['line']", "lines = ast.literal_eval(outputs)['Line']['line']"),
(
"line_type = eval(outputs)['Line']['line_type']",
"line_type = ast.literal_eval(outputs)['Line']['line_type']",
),
(
"endpoints = eval(outputs)['Line']['line_endpoint']",
"endpoints = ast.literal_eval(outputs)['Line']['line_endpoint']",
),
("p0 = eval(line.split(' -- ')[0])", "p0 = ast.literal_eval(line.split(' -- ')[0])"),
("p1 = eval(line.split(' -- ')[-1])", "p1 = ast.literal_eval(line.split(' -- ')[-1])"),
("(x, y) = eval(endpoint.split(': ')[1])", "(x, y) = ast.literal_eval(endpoint.split(': ')[1])"),
)
for old, new in replacements:
text = _replace_exact(text, old, new, 1)
text = _replace_exact(
text,
"images_seq_mask[idx].unsqueeze(-1).cuda()",
"images_seq_mask[idx].unsqueeze(-1).to(inputs_embeds.device)",
1,
)
text = _replace_exact(
text,
" input_ids=input_ids.unsqueeze(0).cuda(),\n",
" input_ids=input_ids.unsqueeze(0).cuda(),\n"
" attention_mask=torch.ones_like(input_ids.unsqueeze(0), device='cuda'),\n",
3,
)
text = _replace_exact(
text,
" eos_token_id=tokenizer.eos_token_id,\n",
" eos_token_id=tokenizer.eos_token_id,\n pad_token_id=tokenizer.eos_token_id,\n",
3,
)
patched_hash = sha256_text(text)
if PATCHED_MODEL_CODE_SHA256 != "__TO_BE_FILLED__" and patched_hash != PATCHED_MODEL_CODE_SHA256:
raise ModelIntegrityError(f"patched model code hash mismatch: {patched_hash}")
return text
def patch_model_file(model_dir: Path) -> str:
source = model_dir / "modeling_unlimitedocr.py"
text = source.read_text(encoding="utf-8")
patched = patch_model_source_text(text)
backup = model_dir / MODEL_UPSTREAM_BACKUP_FILE
if not backup.exists() and sha256_text(text) == UPSTREAM_MODEL_CODE_SHA256:
shutil.copy2(source, backup)
temporary = source.with_suffix(".py.new")
temporary.write_text(patched, encoding="utf-8")
os.replace(temporary, source)
return sha256_text(patched)
def _manifest_payload() -> dict[str, object]:
return {
"schema_version": 1,
"source": MODEL_REPO,
"revision": MODEL_REVISION,
"files": dict(sorted(MODEL_PAYLOAD_SHA256.items())),
"verified_rocm": VERIFIED_ROCM_VERSION,
"verified_torch": VERIFIED_TORCH_VERSION,
}
def _write_manifest(model_dir: Path) -> None:
manifest = model_dir / MODEL_MANIFEST_FILE
descriptor, temporary_name = tempfile.mkstemp(prefix=f".{MODEL_MANIFEST_FILE}.", dir=model_dir)
try:
with os.fdopen(descriptor, "w", encoding="utf-8") as handle:
json.dump(_manifest_payload(), handle, indent=2, sort_keys=True)
handle.write("\n")
handle.flush()
os.fsync(handle.fileno())
os.replace(temporary_name, manifest)
finally:
Path(temporary_name).unlink(missing_ok=True)
def _verify_payload(model_dir: Path, expected_hashes: dict[str, str]) -> dict[str, str]:
if model_dir.is_symlink() or not model_dir.is_dir():
raise ModelIntegrityError(f"model path must be a real directory, not a symlink: {model_dir}")
missing = [name for name in expected_hashes if not _regular_file(model_dir / name)]
if missing:
raise ModelIntegrityError(f"model is incomplete at {model_dir}; missing regular files: {', '.join(missing)}")
weight = model_dir / MODEL_WEIGHT_FILE
weight_bytes = weight.stat().st_size
if weight_bytes != MODEL_WEIGHT_BYTES:
raise ModelIntegrityError(f"weight size mismatch: expected {MODEL_WEIGHT_BYTES}, got {weight_bytes}")
actual_hashes: dict[str, str] = {}
for name, expected in expected_hashes.items():
actual = sha256_file(model_dir / name)
if actual != expected:
raise ModelIntegrityError(f"model file SHA-256 mismatch for {name}: expected {expected}, got {actual}")
actual_hashes[name] = actual
return actual_hashes
def _verify_exact_layout(model_dir: Path) -> None:
if model_dir.is_symlink() or not model_dir.is_dir():
raise ModelIntegrityError(f"model path must be a real directory, not a symlink: {model_dir}")
expected = set(MODEL_PAYLOAD_SHA256) | {MODEL_MANIFEST_FILE}
actual = {entry.name for entry in model_dir.iterdir()}
unexpected = sorted(actual - expected)
missing = sorted(expected - actual)
if unexpected or missing:
details = []
if missing:
details.append(f"missing: {', '.join(missing)}")
if unexpected:
details.append(f"unexpected: {', '.join(unexpected)}")
raise ModelIntegrityError(f"model layout mismatch at {model_dir}; {'; '.join(details)}")
def verify_model(model_dir: Path, *, full_weight_hash: bool = True) -> ModelStatus:
del full_weight_hash # Kept for API compatibility; security verification is always complete.
_verify_exact_layout(model_dir)
hashes = _verify_payload(model_dir, MODEL_PAYLOAD_SHA256)
manifest_path = model_dir / MODEL_MANIFEST_FILE
if not _regular_file(manifest_path):
raise ModelIntegrityError(f"model manifest is missing or not a regular file: {manifest_path}")
try:
manifest = json.loads(manifest_path.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError) as exc:
raise ModelIntegrityError(f"invalid model manifest at {manifest_path}: {exc}") from exc
if manifest != _manifest_payload():
raise ModelIntegrityError(f"model manifest does not match the pinned runtime contract: {manifest_path}")
code_hash = hashes["modeling_unlimitedocr.py"]
weight_hash = hashes[MODEL_WEIGHT_FILE]
weight_bytes = (model_dir / MODEL_WEIGHT_FILE).stat().st_size
return ModelStatus(
model_dir=str(model_dir),
revision=MODEL_REVISION,
prepared=True,
weight_bytes=weight_bytes,
weight_sha256=weight_hash,
model_code_sha256=code_hash,
)
def _remove_snapshot_metadata(model_dir: Path) -> None:
metadata = model_dir / ".cache"
if metadata.is_symlink():
raise ModelIntegrityError(f"refusing symlinked snapshot metadata: {metadata}")
if metadata.exists():
shutil.rmtree(metadata)
def _upgrade_existing_model(model_dir: Path) -> ModelStatus:
allowed = set(MODEL_PAYLOAD_SHA256) | {MODEL_MANIFEST_FILE, ".cache"}
actual = {entry.name for entry in model_dir.iterdir()}
unexpected = sorted(actual - allowed)
if unexpected:
raise ModelIntegrityError(f"refusing unexpected entries in prepared model: {', '.join(unexpected)}")
_verify_payload(model_dir, MODEL_PAYLOAD_SHA256)
_remove_snapshot_metadata(model_dir)
_write_manifest(model_dir)
return verify_model(model_dir)
def _lock_handle(path: Path) -> BinaryIO:
descriptor = os.open(path, os.O_RDWR | os.O_CREAT | os.O_CLOEXEC | os.O_NOFOLLOW, 0o600)
handle = os.fdopen(descriptor, "a+b")
flock(handle.fileno(), LOCK_EX)
return handle
def prepare_model(model_dir: Path, cache_dir: Path, *, dry_run: bool = False) -> ModelStatus:
model_dir = model_dir.expanduser().absolute()
cache_dir = cache_dir.expanduser().resolve()
if dry_run:
return ModelStatus(
model_dir=str(model_dir),
revision=MODEL_REVISION,
prepared=False,
weight_bytes=MODEL_WEIGHT_BYTES,
weight_sha256=MODEL_WEIGHT_SHA256,
model_code_sha256=PATCHED_MODEL_CODE_SHA256,
)
if model_dir.is_symlink():
raise ModelIntegrityError(f"model destination must not be a symlink: {model_dir}")
model_dir.parent.mkdir(parents=True, exist_ok=True)
cache_dir.mkdir(parents=True, exist_ok=True)
try:
from huggingface_hub import snapshot_download
except ImportError as exc:
raise ModelIntegrityError("huggingface-hub is missing; install the package dependencies") from exc
lock_path = model_dir.parent / f".{model_dir.name}.prepare.lock"
with _lock_handle(lock_path):
if model_dir.exists():
if not model_dir.is_dir():
raise ModelIntegrityError(f"model destination exists and is not a directory: {model_dir}")
status = _upgrade_existing_model(model_dir)
return ModelStatus(**{**asdict(status), "reused": True})
partial = Path(tempfile.mkdtemp(prefix=f".{model_dir.name}.partial-", dir=model_dir.parent))
try:
snapshot_download(
repo_id=MODEL_REPO,
revision=MODEL_REVISION,
local_dir=partial,
cache_dir=cache_dir / "huggingface",
allow_patterns=sorted(MODEL_DOWNLOAD_SHA256),
)
_verify_payload(partial, MODEL_DOWNLOAD_SHA256)
patch_model_file(partial)
_remove_snapshot_metadata(partial)
_write_manifest(partial)
status = verify_model(partial)
partial.rename(model_dir)
except Exception:
shutil.rmtree(partial, ignore_errors=True)
raise
return ModelStatus(**{**asdict(status), "model_dir": str(model_dir)})