"""Official PNG codec compatible weight reconstruction adapters.""" from __future__ import annotations import hashlib import json from collections.abc import Mapping, Sequence from pathlib import Path import numpy as np from PIL import Image def decode_weight_image(image_path: str | Path, manifest_path: str | Path) -> tuple[np.ndarray, list[str]]: """公式manifest順でRGBのR/Gをfp16 bit patternとしてfloat32へ復元する。""" metadata = json.loads(Path(manifest_path).read_text()) pixels = np.asarray(Image.open(image_path).convert("RGB"), dtype=np.uint8).reshape(-1, 3) total = int(metadata["total_parameters"]) if len(pixels) < total: raise ValueError(f"weight image has {len(pixels)} pixels, needs {total}") bits = (pixels[:total, 0].astype(np.uint16) << 8) | pixels[:total, 1].astype(np.uint16) return bits.view(np.float16).astype(np.float32), [str(item["name"]) for item in metadata["params"]] def flatten_manifest(arrays: Mapping[str, np.ndarray]) -> tuple[np.ndarray, list[tuple[str, int, int]]]: """mappingの挿入順を保持し、float16 quantization後のfloat32 flatを返す。""" parts: list[np.ndarray] = [] manifest: list[tuple[str, int, int]] = [] offset = 0 for name, value in arrays.items(): part = np.asarray(value).reshape(-1).astype(np.float16).astype(np.float32) parts.append(part) manifest.append((name, offset, offset + len(part))) offset += len(part) return np.concatenate(parts) if parts else np.empty(0, dtype=np.float32), manifest def layer_offsets(names: Sequence[str], sizes: Sequence[int]) -> dict[str, tuple[int, int]]: """layer名とnumelからmetric用offsetを構築する。""" offsets: dict[str, tuple[int, int]] = {} offset = 0 for name, size in zip(names, sizes): offsets[name] = (offset, offset + int(size)) offset += int(size) return offsets def sha256_file(path: str | Path) -> str: """ファイルのSHA-256を返す。""" digest = hashlib.sha256() with Path(path).open("rb") as stream: for chunk in iter(lambda: stream.read(1024 * 1024), b""): digest.update(chunk) return digest.hexdigest() def reconstruct_dit(image_path: str | Path, manifest_path: str | Path, device: str = "cpu", config_path: str | Path | None = None): """model_png.jsonのname/shape順でDiTを再構築し、copy後の全parameterをfinite gateする。""" import torch metadata = json.loads(Path(manifest_path).read_text(encoding="utf-8")) config = json.loads(Path(config_path).read_text(encoding="utf-8"))["dit"] if config_path is not None else metadata["cfg"] from dit import DiT model = DiT(dim=int(config["dim"]), depth=int(config["depth"]), heads=int(config["heads"])) named = dict(model.named_parameters()) pixels = np.asarray(Image.open(image_path).convert("RGB"), dtype=np.uint8).reshape(-1, 3) total = int(metadata["total_parameters"]) if pixels.shape[0] < total: raise ValueError("weight image is shorter than manifest") bits = (pixels[:total, 0].astype(np.uint16) << 8) | pixels[:total, 1].astype(np.uint16) flat = bits.view(np.float16) offset = 0 with torch.no_grad(): for item in metadata["params"]: name = str(item["name"]) shape = tuple(int(value) for value in item["shape"]) numel = int(item["numel"]) if name not in named or tuple(named[name].shape) != shape or numel != int(np.prod(shape)): raise ValueError(f"manifest mismatch for {name}") named[name].copy_(torch.from_numpy(flat[offset:offset + numel].copy()).view(shape).float()) offset += numel if offset != total or any(not bool(torch.isfinite(parameter).all()) for parameter in model.parameters()): raise ValueError("finite gate failed after model parameter copy") return model.to(device).eval()