| """公式DiTの再現可能な推論と、テスト用runtime依存性注入。""" |
|
|
| from __future__ import annotations |
|
|
| import math |
| import os |
| from contextlib import nullcontext |
| from typing import Any |
|
|
| os.environ.setdefault("TOKENIZERS_PARALLELISM", "false") |
|
|
|
|
| def resolve_device(requested: str = "auto") -> str: |
| """cpu/mps/cudaを環境に応じて解決する。""" |
| import torch |
| if requested != "auto": |
| if requested == "cuda" and not torch.cuda.is_available(): |
| raise RuntimeError("CUDA was requested but is unavailable") |
| if requested == "mps" and not torch.backends.mps.is_available(): |
| raise RuntimeError("MPS was requested but is unavailable") |
| return requested |
| if torch.cuda.is_available(): |
| return "cuda" |
| if torch.backends.mps.is_available(): |
| return "mps" |
| return "cpu" |
|
|
|
|
| def autocast_context(device: str) -> Any: |
| """GPUではautocast、CPUではnullcontextを返す。""" |
| import torch |
| if device == "cuda": |
| return torch.autocast(device_type="cuda", dtype=torch.bfloat16) |
| if device == "mps": |
| return torch.autocast(device_type="mps", dtype=torch.float16) |
| return nullcontext() |
|
|
|
|
| def finite_state_dict(state_dict: dict[str, Any]) -> bool: |
| """state dict全tensorのfinite gate。""" |
| return all(bool(value.isfinite().all()) for value in state_dict.values() if hasattr(value, "isfinite")) |
|
|
|
|
| def euler_rectified_flow(model: Any, conditional: Any, unconditional: Any, latent: Any, steps: int, cfg: float, device: str) -> Any: |
| """公式と同じ50-step Euler rectified-flow/CFGを実行する。""" |
| if steps < 1: |
| raise ValueError("steps must be positive") |
| dt = 1.0 / steps |
| x = latent.clone() |
| for index in range(steps): |
| t = x.new_full((x.shape[0],), index * dt, device=device) |
| with autocast_context(device): |
| vc = model(x, t, conditional[0], conditional[1]) |
| vu = model(x, t, unconditional[0], unconditional[1]) |
| x = x + (vu + cfg * (vc - vu)).float() * dt |
| return x |
|
|
|
|
| def generate_with_runtime(runtime: Any, model: Any, prompt: str, seed: int, steps: int, cfg: float, device: str) -> tuple[Any, float | None]: |
| """text encode→seed固定latent→Euler→VAE decode→CLIPをruntime経由で行う。""" |
| import torch |
| with torch.inference_mode(): |
| conditional = runtime.encode([prompt], device) |
| unconditional = runtime.encode([""], device) |
| latent = runtime.initial_latent(seed, device) |
| result = euler_rectified_flow(model, conditional, unconditional, latent, steps, cfg, device) |
| image = runtime.decode(result, device) |
| if not bool(torch.isfinite(image).all()): |
| raise FloatingPointError("nonfinite generated image tensor") |
| score = runtime.clip_score(image, prompt, device) |
| if score is not None and not math.isfinite(float(score)): |
| raise FloatingPointError("nonfinite CLIP score") |
| return image, float(score) if score is not None else None |
|
|
|
|
| def clip_score_once(metric: Any, image: Any, prompt: str, target: str, torch_module: Any) -> float: |
| """stateful torchmetrics CLIPScoreを一サンプル単位で隔離して評価する。""" |
| metric.reset() |
| try: |
| metric((image.unsqueeze(0) * 255).to(torch_module.uint8), [prompt]) |
| return float(metric.compute().item()) |
| finally: |
| metric.reset() |
|
|
|
|
| def build_official_runtime(args: Any, device: str) -> tuple[Any, Any]: |
| """公式依存を一度だけロードし、model loaderとruntimeを返す。""" |
| import torch |
| from diffusers import AutoencoderKL |
| from torchmetrics.multimodal.clip_score import CLIPScore |
| from transformers import CLIPTextModel, CLIPTokenizer |
| vae = AutoencoderKL.from_pretrained(args.vae).to(device).half().eval() |
| tokenizer = CLIPTokenizer.from_pretrained(args.clip) |
| text = CLIPTextModel.from_pretrained(args.clip).to(device).half().eval() |
| clip_metric = CLIPScore(model_name_or_path=args.clip).to(device).eval() |
|
|
| class Runtime: |
| def encode(self, strings: list[str], target: str) -> tuple[Any, Any]: |
| tokens = tokenizer(strings, padding="max_length", max_length=40, truncation=True, return_tensors="pt").to(target) |
| output = text(**tokens) |
| return output.last_hidden_state.float(), output.pooler_output.float() |
|
|
| def initial_latent(self, seed: int, target: str) -> Any: |
| generator = torch.Generator(device=target).manual_seed(seed) |
| return torch.randn((1, 4, 32, 32), generator=generator, device=target) |
|
|
| def decode(self, latent: Any, target: str) -> Any: |
| return ((vae.decode((latent / 0.18215).half()).sample.float().clamp(-1, 1) + 1) / 2)[0] |
|
|
| def clip_score(self, image: Any, prompt: str, target: str) -> float | None: |
| return clip_score_once(clip_metric, image, prompt, target, torch) |
|
|
| from pixelmodel_robustness.codec import reconstruct_dit |
|
|
| def loader(weight_path: str) -> Any: |
| model = reconstruct_dit(weight_path, args.manifest, device, getattr(args, "config", None)) |
| if not finite_state_dict(dict(model.named_parameters())): |
| raise ValueError("finite gate failed after PNG model reconstruction") |
| return model |
|
|
| return loader, Runtime() |
|
|