File size: 5,145 Bytes
85d0fdc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Minimal vendored FLUX.2 klein adapter for the HF Space (from FD-Loss train/flux2_adapter.py).

Differences vs the training adapter: flux2 is pip-installed (no sys.path hack), weights come from a
provided state_dict (the epfl-vita/flux2-klein-1step-rdm model.safetensors, "model."-prefixed keys OK),
no grad-checkpoint / compile / disc-feature paths.
"""
from __future__ import annotations

import torch
import torch.nn as nn

FLUX2_VAE_DOWNSAMPLE = 16
FLUX2_LATENT_CHANNELS = 128


class Flux2AdapterModel(nn.Module):
    """FLUX.2 klein-4B MM-DiT + 1-step (or N-step) flow-matching Euler sampler."""

    def __init__(self, state_dict: dict, image_resolution: int = 512,
                 param_dtype: torch.dtype = torch.bfloat16, guidance: float = 1.0):
        super().__init__()
        import flux2.model as _fm
        from flux2 import sampling as _sampling

        self.image_resolution = int(image_resolution)
        self.guidance = float(guidance)
        assert self.image_resolution % FLUX2_VAE_DOWNSAMPLE == 0

        params = _fm.Klein4BParams()
        assert params.in_channels == FLUX2_LATENT_CHANNELS
        with torch.device("meta"):
            model = _fm.Flux2(params).to(torch.bfloat16)
        if all(k.startswith("model.") for k in list(state_dict.keys())[:8]):
            state_dict = {k[len("model."):]: v for k, v in state_dict.items()}
        model.load_state_dict(state_dict, strict=True, assign=True)
        self.model = model.to(dtype=param_dtype)

        self._batched_prc_img = _sampling.batched_prc_img
        self._batched_prc_txt = _sampling.batched_prc_txt
        self._get_schedule = _sampling.get_schedule
        self._timestep_embedding = _fm.timestep_embedding

        self.in_channels = FLUX2_LATENT_CHANNELS
        self.input_size = self.image_resolution // FLUX2_VAE_DOWNSAMPLE

    @property
    def device(self) -> torch.device:
        return next(self.model.parameters()).device

    def _run_dit(self, x, x_ids, ctx, ctx_ids, t_vec):
        m = self.model
        num_txt_tokens = ctx.shape[1]
        vec = m.time_in(self._timestep_embedding(t_vec, 256))
        if getattr(m, "use_guidance_embed", False):   # klein: False -> skipped
            guid = torch.full((x.shape[0],), self.guidance, dtype=x.dtype, device=x.device)
            vec = vec + m.guidance_in(self._timestep_embedding(guid, 256))
        mod_img = m.double_stream_modulation_img(vec)
        mod_txt = m.double_stream_modulation_txt(vec)
        single_mod, _ = m.single_stream_modulation(vec)
        img = m.img_in(x)
        txt = m.txt_in(ctx)
        pe_x = m.pe_embedder(x_ids)
        pe_ctx = m.pe_embedder(ctx_ids)
        for block in m.double_blocks:
            img, txt, _ = block.forward_kv_extract(img, txt, pe_x, pe_ctx, mod_img, mod_txt, 0)
        img = torch.cat((txt, img), dim=1)
        pe = torch.cat((pe_ctx, pe_x), dim=2)
        for block in m.single_blocks:
            img, _ = block.forward_kv_extract(img, pe, single_mod, num_txt_tokens, 0)
        img = img[:, num_txt_tokens:, ...]
        return m.final_layer(img, vec)

    def sample_images_with_grad(self, noise: torch.Tensor, condition: torch.Tensor,
                                sampling_args: dict) -> torch.Tensor:
        """noise (B,128,H,W) + Qwen3 ctx (B,L,7680) -> normalized latents (B,128,H,W)."""
        B = noise.shape[0]
        ctx = condition.to(device=noise.device, dtype=noise.dtype)
        x, x_ids = self._batched_prc_img(noise)
        ctx, ctx_ids = self._batched_prc_txt(ctx)
        H, W = noise.shape[-2], noise.shape[-1]
        num_steps = int(sampling_args.get("num_steps", 1))
        timesteps = self._get_schedule(num_steps, x.shape[1])
        for t_curr, t_prev in zip(timesteps[:-1], timesteps[1:]):
            t_vec = torch.full((B,), t_curr, dtype=x.dtype, device=x.device)
            pred = self._run_dit(x, x_ids, ctx, ctx_ids, t_vec)
            x = x + (t_prev - t_curr) * pred
        from einops import rearrange
        return rearrange(x, "b (h w) c -> b c h w", h=H, w=W)


class Flux2VAETokenizer(nn.Module):
    """Native FLUX.2 AutoEncoder (BFL ae.safetensors). decode -> [-1,1]; detokenize -> [0,1]."""

    def __init__(self, ae_path: str, device="cpu", torch_dtype: torch.dtype = torch.bfloat16):
        super().__init__()
        from flux2.autoencoder import AutoEncoder, AutoEncoderParams
        from safetensors.torch import load_file as load_sft

        with torch.device("meta"):
            ae = AutoEncoder(AutoEncoderParams())
        sd = load_sft(ae_path, device="cpu")
        ae.load_state_dict(sd, strict=True, assign=True)
        ae = ae.to(device=device, dtype=torch_dtype)
        for p in ae.parameters():
            p.requires_grad = False
        ae.eval()
        self.vae = ae

    def denormalize_z(self, z):
        return z  # AE.decode applies inv-normalize internally

    def decode(self, z):
        z = z.to(dtype=next(self.vae.parameters()).dtype)
        return self.vae.decode(z)

    @torch.inference_mode()
    def detokenize(self, z):
        return torch.clamp(self.decode(self.denormalize_z(z)) * 0.5 + 0.5, 0.0, 1.0)