File size: 5,329 Bytes
9f08d74
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
120
121
122
"""公式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()